Compare commits
174 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| c051929f8c | |||
| 4814ec2363 | |||
| 98b9d612fa | |||
| 874873541d | |||
| ef9ac76f06 | |||
| ca9db195de | |||
| c1e16f2163 | |||
| 61d32fa00f | |||
| cddb3f8ccf | |||
| 87a938fe58 | |||
| 9157694412 | |||
| aa34314175 | |||
| c7fc386d0f | |||
| 6bd814e346 | |||
| aa25852091 | |||
| c2740cd282 | |||
| 23875bb3cc | |||
| ee1eab8408 | |||
| 2254ba7496 | |||
| 4866ad08b1 | |||
| 97ebc19313 | |||
| 946ac6f66a | |||
| 90ddee14c2 | |||
| f17f8ae8c8 | |||
| d713bfc5f9 | |||
| 27fe8c24ec | |||
| eef1e4e8c6 | |||
| a7a5905874 | |||
| 021537de56 | |||
| fdfd75790d | |||
| 3d1f8d9cf7 | |||
| eec37ab710 | |||
| 5f499ec1b0 | |||
| 2b5223097f | |||
| 1efbc289ba | |||
| 3c57c8b78b | |||
| 9c4500f9cb | |||
| 1e2c5a68cf | |||
| dc6991e6ef | |||
| cdf233378c | |||
| 23769e6ad4 | |||
| 9f8058223a | |||
| b483a34517 | |||
| a8fbca9ef5 | |||
| de8c47c81c | |||
| b11089896a | |||
| 16549709e2 | |||
| 68977b73be | |||
| 3da2bc93cb | |||
| d9632a3412 | |||
| 03d22c4d06 | |||
| 19242b4152 | |||
| ceaed9af66 | |||
| e9364fa70f | |||
| 2bec205a30 | |||
| cbd3436ff7 | |||
| 57ff236f2d | |||
| 6085859874 | |||
| d413fb84a5 | |||
| c22b047b8c | |||
| 090edf7ef6 | |||
| cbaebcf649 | |||
| 4a0d38384f | |||
| ea006b68c2 | |||
| 202533fbb1 | |||
| 0952aa8217 | |||
| 787e8844bc | |||
| fe98f966d6 | |||
| 79ad3b0715 | |||
| f728011b2a | |||
| 569876538a | |||
| d2b1c132d1 | |||
| 8926152fca | |||
| 2682e0268c | |||
| f13a63dc2f | |||
| 4a7f24348c | |||
| 0fdbef578f | |||
| 29a12bb102 | |||
| 270764ff0f | |||
| 0e6521e433 | |||
| e20d94d6ba | |||
| f4802ece4d | |||
| 9800e37cd6 | |||
| 84f30b07c4 | |||
| fba5e7c7be | |||
| 1e7bd0a540 | |||
| 0a0af4e02a | |||
| 3aa56dcfc3 | |||
| 7e3c701ea6 | |||
| e672bdde54 | |||
| c7c4cb45a7 | |||
| 6e4275a510 | |||
| 3ef60be623 | |||
| a3d01dd0a5 | |||
| 9bd5fc5c68 | |||
| b6e71e32f5 | |||
| 9ccaae04db | |||
| 9f90624aa6 | |||
| 62c1fb3836 | |||
| 569c20cf63 | |||
| f658b71079 | |||
| 3e99e7f197 | |||
| c2c983a01e | |||
| e46b4f9249 | |||
| 8eb851793d | |||
| 8e5e815ac9 | |||
| 143a254b0c | |||
| 5deee8c65c | |||
| 62d1bdc462 | |||
| 0b35ae3bf0 | |||
| b55300ff6f | |||
| 314ba3aee4 | |||
| 18e4a89573 | |||
| 29943ac239 | |||
| 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 | |||
| 9e88acaa36 |
@@ -0,0 +1,3 @@
|
|||||||
|
{
|
||||||
|
"fingerprint": "fdea8a74bb4c7449c01c4bd61646c895b10ede78"
|
||||||
|
}
|
||||||
@@ -0,0 +1,36 @@
|
|||||||
|
# Skill Registry — headquarter
|
||||||
|
|
||||||
|
<!-- Auto-generated by gentle-pi extensions/skill-registry.ts. Run /skill-registry:refresh to regenerate. -->
|
||||||
|
|
||||||
|
Last updated: 2026-05-28
|
||||||
|
|
||||||
|
## Sources scanned
|
||||||
|
|
||||||
|
- .opencode/skills
|
||||||
|
- .claude/skills
|
||||||
|
- /home/alex/.config/opencode/skills
|
||||||
|
|
||||||
|
## Contract
|
||||||
|
|
||||||
|
**Delegator use only.** This registry is an index, not a summary. Any agent that launches subagents reads it to select relevant skills, then passes exact `SKILL.md` paths for the subagent to read before work.
|
||||||
|
|
||||||
|
`SKILL.md` remains the source of truth. Do not inject generated summaries or compact rules by default; pass paths so subagents load the full runtime contract and preserve author intent.
|
||||||
|
|
||||||
|
## Skills
|
||||||
|
|
||||||
|
| Skill | Trigger / description | Scope | Path |
|
||||||
|
| --- | --- | --- | --- |
|
||||||
|
| `auto-commit` | Use when you are making multiple edits or completing significant work in a git repository to automatically create commits | user | `/home/alex/.config/opencode/skills/auto-commit/SKILL.md` |
|
||||||
|
| `openspec` | Use OpenSpec as the source of truth for planning, implementation, verification, and archive discipline. | user | `/home/alex/.config/opencode/skills/openspec/SKILL.md` |
|
||||||
|
| `openspec-apply-change` | Implement tasks from an OpenSpec change. Use when the user wants to start implementing, continue implementation, or work through tasks. | project | `/home/alex/projects/headquarter/.opencode/skills/openspec-apply-change/SKILL.md` |
|
||||||
|
| `openspec-archive-change` | Archive a completed change in the experimental workflow. Use when the user wants to finalize and archive a change after implementation is complete. | project | `/home/alex/projects/headquarter/.opencode/skills/openspec-archive-change/SKILL.md` |
|
||||||
|
| `openspec-explore` | Enter explore mode - a thinking partner for exploring ideas, investigating problems, and clarifying requirements. Use when the user wants to think through something before or during a change. | project | `/home/alex/projects/headquarter/.opencode/skills/openspec-explore/SKILL.md` |
|
||||||
|
| `openspec-propose` | Propose a new change with all artifacts generated in one step. Use when the user wants to quickly describe what they want to build and get a complete proposal with design, specs, and tasks ready for implementation. | project | `/home/alex/projects/headquarter/.opencode/skills/openspec-propose/SKILL.md` |
|
||||||
|
| `sift-backlog` | 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. | project | `/home/alex/projects/headquarter/.claude/skills/sift-backlog/SKILL.md` |
|
||||||
|
|
||||||
|
## Loading protocol
|
||||||
|
|
||||||
|
1. Match task context and target files against the `Trigger / description` column.
|
||||||
|
2. Pass only the matching `Path` values to the subagent under `## Skills to load before work`.
|
||||||
|
3. Instruct the subagent to read those exact `SKILL.md` files before reading, writing, reviewing, testing, or creating artifacts.
|
||||||
|
4. If no matching skill exists, proceed without project skill injection and report `skill_resolution: none`.
|
||||||
@@ -49,3 +49,8 @@ apps/web/dist/
|
|||||||
.DS_Store
|
.DS_Store
|
||||||
Thumbs.db
|
Thumbs.db
|
||||||
/.stoneforge/.worktrees/
|
/.stoneforge/.worktrees/
|
||||||
|
# Pi / agent cache
|
||||||
|
.pi/
|
||||||
|
.atl/
|
||||||
|
.sisyphus/
|
||||||
|
.pi-lens/
|
||||||
|
|||||||
@@ -0,0 +1 @@
|
|||||||
|
{}
|
||||||
@@ -0,0 +1,10 @@
|
|||||||
|
{
|
||||||
|
"sessionID": "ses_1da2608b1ffergOzow3NQt1mGr",
|
||||||
|
"updatedAt": "2026-05-15T23:50:42.832Z",
|
||||||
|
"sources": {
|
||||||
|
"background-task": {
|
||||||
|
"state": "idle",
|
||||||
|
"updatedAt": "2026-05-15T23:50:42.832Z"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -4,6 +4,10 @@
|
|||||||
|
|
||||||
OpenSpec is the source of truth. Superpowers is the default workflow. Keep changes small, scoped, and verified.
|
OpenSpec is the source of truth. Superpowers is the default workflow. Keep changes small, scoped, and verified.
|
||||||
|
|
||||||
|
## Communication
|
||||||
|
|
||||||
|
All agent output, code comments, commit messages, documentation, and artifacts must be in **English** unless the user explicitly requests another language.
|
||||||
|
|
||||||
## Priority order
|
## Priority order
|
||||||
|
|
||||||
1. Current user instruction
|
1. Current user instruction
|
||||||
|
|||||||
+568
@@ -0,0 +1,568 @@
|
|||||||
|
{
|
||||||
|
"version": "v2",
|
||||||
|
"timestamp": 1779889907001,
|
||||||
|
"ruleHash": "fd9b2b15f2ac8993",
|
||||||
|
"queries": [
|
||||||
|
{
|
||||||
|
"id": "bare-except",
|
||||||
|
"name": "Bare Except Clause",
|
||||||
|
"severity": "warning",
|
||||||
|
"language": "python",
|
||||||
|
"message": "Bare 'except:' clause — catches SystemExit, KeyboardInterrupt",
|
||||||
|
"query": " (except_clause\n \"except\") @CLAUSE",
|
||||||
|
"metavars": [
|
||||||
|
"CLAUSE"
|
||||||
|
],
|
||||||
|
"post_filter": "bare_except_only",
|
||||||
|
"defect_class": "silent-error",
|
||||||
|
"inline_tier": "blocking",
|
||||||
|
"filePath": "/home/alex/.npm-global/lib/node_modules/pi-lens/rules/tree-sitter-queries/python/bare-except.yml"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": "eval-exec",
|
||||||
|
"name": "Eval/Exec Usage",
|
||||||
|
"severity": "warning",
|
||||||
|
"language": "python",
|
||||||
|
"message": "{{FUNC}}() detected — security risk, code injection vulnerability",
|
||||||
|
"query": " (call\n function: (identifier) @FUNC\n (#match? @FUNC \"^(eval|exec)$\")\n arguments: (argument_list) @ARGS)",
|
||||||
|
"metavars": [
|
||||||
|
"FUNC",
|
||||||
|
"ARGS"
|
||||||
|
],
|
||||||
|
"defect_class": "injection",
|
||||||
|
"inline_tier": "blocking",
|
||||||
|
"filePath": "/home/alex/.npm-global/lib/node_modules/pi-lens/rules/tree-sitter-queries/python/eval-exec.yml"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": "exit-signature-check",
|
||||||
|
"name": "__exit__ Missing Parameters",
|
||||||
|
"severity": "error",
|
||||||
|
"language": "python",
|
||||||
|
"message": "__exit__ should accept type, value, and traceback arguments",
|
||||||
|
"query": " (function_definition\n name: (identifier) @NAME (#eq? @NAME \"__exit__\")\n parameters: (parameters\n (_) @SELF\n . (_) @PARAM1?\n . (_) @PARAM2?\n . (_) @PARAM3?))",
|
||||||
|
"metavars": [
|
||||||
|
"NAME",
|
||||||
|
"SELF",
|
||||||
|
"PARAM1",
|
||||||
|
"PARAM2",
|
||||||
|
"PARAM3"
|
||||||
|
],
|
||||||
|
"post_filter": "exit_params_insufficient",
|
||||||
|
"defect_class": "correctness",
|
||||||
|
"inline_tier": "blocking",
|
||||||
|
"filePath": "/home/alex/.npm-global/lib/node_modules/pi-lens/rules/tree-sitter-queries/python/exit-signature-check.yml"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": "in-operator-unsupported",
|
||||||
|
"name": "In and Not In Operators Should Be Used on Valid Objects",
|
||||||
|
"severity": "warning",
|
||||||
|
"language": "python",
|
||||||
|
"message": "'in' operator used on object that may not support containment",
|
||||||
|
"query": " (comparison_operator\n (identifier) @OBJ\n \"in\"\n (identifier) @TARGET)\n (comparison_operator\n (identifier) @OBJ\n \"not\"\n \"in\"\n (identifier) @TARGET)",
|
||||||
|
"metavars": [
|
||||||
|
"OBJ",
|
||||||
|
"TARGET"
|
||||||
|
],
|
||||||
|
"post_filter": "check_in_operator_types",
|
||||||
|
"defect_class": "correctness",
|
||||||
|
"inline_tier": "warning",
|
||||||
|
"filePath": "/home/alex/.npm-global/lib/node_modules/pi-lens/rules/tree-sitter-queries/python/in-operator-unsupported.yml"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": "is-vs-equals",
|
||||||
|
"name": "Is vs Equals for Literals",
|
||||||
|
"severity": "warning",
|
||||||
|
"language": "python",
|
||||||
|
"message": "Using 'is' with literal — use '==' for value comparison",
|
||||||
|
"query": " (comparison_operator\n (identifier)\n (\"is\")\n (string) @LITERAL)\n (comparison_operator\n (identifier)\n (\"is not\")\n (string) @LITERAL)\n (comparison_operator\n (identifier)\n (\"is\")\n (integer) @LITERAL)\n (comparison_operator\n (identifier)\n (\"is not\")\n (integer) @LITERAL)",
|
||||||
|
"metavars": [
|
||||||
|
"LITERAL"
|
||||||
|
],
|
||||||
|
"defect_class": "correctness",
|
||||||
|
"inline_tier": "blocking",
|
||||||
|
"filePath": "/home/alex/.npm-global/lib/node_modules/pi-lens/rules/tree-sitter-queries/python/is-vs-equals.yml"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": "iter-return-iterator",
|
||||||
|
"name": "__iter__ Should Return Iterator",
|
||||||
|
"severity": "warning",
|
||||||
|
"language": "python",
|
||||||
|
"message": "__iter__ should return an iterator (object with __next__ method)",
|
||||||
|
"query": " (function_definition\n name: (identifier) @NAME (#eq? @NAME \"__iter__\")\n body: (block\n (return_statement) @RETURN))",
|
||||||
|
"metavars": [
|
||||||
|
"NAME",
|
||||||
|
"RETURN"
|
||||||
|
],
|
||||||
|
"defect_class": "correctness",
|
||||||
|
"inline_tier": "blocking",
|
||||||
|
"filePath": "/home/alex/.npm-global/lib/node_modules/pi-lens/rules/tree-sitter-queries/python/iter-return-iterator.yml"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": "mutable-default-arg",
|
||||||
|
"name": "Mutable Default Argument",
|
||||||
|
"severity": "warning",
|
||||||
|
"language": "python",
|
||||||
|
"message": "Mutable default argument — list/dict/set as default value",
|
||||||
|
"query": " (function_definition\n (parameters\n (default_parameter\n (identifier) @PARAM\n [(list) (dictionary) (set)] @MUTABLE)))",
|
||||||
|
"metavars": [
|
||||||
|
"PARAM",
|
||||||
|
"MUTABLE"
|
||||||
|
],
|
||||||
|
"defect_class": "correctness",
|
||||||
|
"inline_tier": "blocking",
|
||||||
|
"filePath": "/home/alex/.npm-global/lib/node_modules/pi-lens/rules/tree-sitter-queries/python/mutable-default-arg.yml"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": "no-super-torchscript",
|
||||||
|
"name": "super Should Not Be Used in TorchScript Methods",
|
||||||
|
"severity": "error",
|
||||||
|
"language": "python",
|
||||||
|
"message": "super() calls should not be used in TorchScript methods",
|
||||||
|
"query": " (function_definition\n (decorator\n (call\n function: (identifier) @DEC (#match? @DEC \"^(torch\\.jit\\.script|jit\\.script)$\")))\n body: (block\n (call\n function: (identifier) @FUNC (#eq? @FUNC \"super\")) @CALL))",
|
||||||
|
"metavars": [
|
||||||
|
"DEC",
|
||||||
|
"FUNC",
|
||||||
|
"CALL"
|
||||||
|
],
|
||||||
|
"defect_class": "correctness",
|
||||||
|
"inline_tier": "blocking",
|
||||||
|
"filePath": "/home/alex/.npm-global/lib/node_modules/pi-lens/rules/tree-sitter-queries/python/no-super-torchscript.yml"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": "notimplemented-boolean-context",
|
||||||
|
"name": "NotImplemented in Boolean Context",
|
||||||
|
"severity": "error",
|
||||||
|
"language": "python",
|
||||||
|
"message": "NotImplemented should not be used in boolean contexts",
|
||||||
|
"query": " (if_statement\n condition: (identifier) @COND (#eq? @COND \"NotImplemented\"))\n (while_statement\n condition: (identifier) @COND (#eq? @COND \"NotImplemented\"))\n (binary_operator\n (identifier) @COND (#eq? @COND \"NotImplemented\")\n (\"and\" | \"or\"))\n (boolean_operator\n (identifier) @COND (#eq? @COND \"NotImplemented\"))\n (unary_operator\n operator: (\"not\")\n argument: (identifier) @COND (#eq? @COND \"NotImplemented\"))",
|
||||||
|
"metavars": [
|
||||||
|
"COND"
|
||||||
|
],
|
||||||
|
"defect_class": "correctness",
|
||||||
|
"inline_tier": "blocking",
|
||||||
|
"filePath": "/home/alex/.npm-global/lib/node_modules/pi-lens/rules/tree-sitter-queries/python/notimplemented-boolean-context.yml"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": "python-assert-production",
|
||||||
|
"name": "Assert in Production Code",
|
||||||
|
"severity": "warning",
|
||||||
|
"language": "python",
|
||||||
|
"message": "assert statement stripped by Python -O flag — use explicit checks with exceptions in production code",
|
||||||
|
"query": " (assert_statement) @ASSERT",
|
||||||
|
"metavars": [
|
||||||
|
"ASSERT"
|
||||||
|
],
|
||||||
|
"defect_class": "correctness",
|
||||||
|
"inline_tier": "warning",
|
||||||
|
"filePath": "/home/alex/.npm-global/lib/node_modules/pi-lens/rules/tree-sitter-queries/python/python-assert-production.yml"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": "python-command-injection",
|
||||||
|
"name": "Command Injection Sink",
|
||||||
|
"severity": "error",
|
||||||
|
"language": "python",
|
||||||
|
"message": "Potential command injection sink — avoid shell execution with dynamic input",
|
||||||
|
"query": " (call\n function: (attribute\n object: (identifier) @MOD\n attribute: (identifier) @FN)\n arguments: (argument_list) @ARGS\n (#eq? @MOD \"os\")\n (#match? @FN \"^(system|popen)$\"))\n\n (call\n function: (attribute\n object: (identifier) @MOD\n attribute: (identifier) @FN)\n arguments: (argument_list\n (keyword_argument\n name: (identifier) @KW\n value: (true)))\n (#eq? @MOD \"subprocess\")\n (#match? @FN \"^(run|Popen|call|check_output|check_call)$\")\n (#eq? @KW \"shell\"))",
|
||||||
|
"metavars": [
|
||||||
|
"MOD",
|
||||||
|
"FN",
|
||||||
|
"ARGS",
|
||||||
|
"KW"
|
||||||
|
],
|
||||||
|
"post_filter": "py_command_injection_sink",
|
||||||
|
"defect_class": "injection",
|
||||||
|
"inline_tier": "blocking",
|
||||||
|
"filePath": "/home/alex/.npm-global/lib/node_modules/pi-lens/rules/tree-sitter-queries/python/python-command-injection.yml"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": "python-cross-language-method",
|
||||||
|
"name": "Cross-Language Method Leakage",
|
||||||
|
"severity": "warning",
|
||||||
|
"language": "python",
|
||||||
|
"message": "'{METHOD}' is not a Python method — likely a {LANG} idiom leaking in",
|
||||||
|
"query": " (call\n function: (attribute\n object: (_) @OBJ\n attribute: (identifier) @METHOD)\n (#match? @METHOD \"^(push|forEach|indexOf|charAt|substring|hasOwnProperty|unshift|flatMap|padStart|padEnd|trimStart|trimEnd|equals|isEmpty|println|printf|getClass|hashCode|toCharArray|getBytes|compareTo|equalsIgnoreCase|startsWith|endsWith|each|collect|select|reject|detect|inject|chomp|chop|gsub|upcase|downcase|present|blank|Add|Contains|ToLower|ToUpper|Trim|Substring|WriteLine|ReadLine|TryParse|forEach|includes|assign|freeze|splice|unshift|shift|flatMap)$\"))",
|
||||||
|
"metavars": [
|
||||||
|
"OBJ",
|
||||||
|
"METHOD"
|
||||||
|
],
|
||||||
|
"post_filter": "match_captures",
|
||||||
|
"post_filter_params": {
|
||||||
|
"METHOD": "^(push|forEach|indexOf|charAt|substring|hasOwnProperty|unshift|flatMap|padStart|padEnd|trimStart|trimEnd|equals|isEmpty|println|printf|getClass|hashCode|toCharArray|getBytes|compareTo|equalsIgnoreCase|startsWith|endsWith|each|collect|select|reject|detect|inject|chomp|chop|gsub|upcase|downcase|present|blank|Add|Contains|ToLower|ToUpper|Trim|Substring|WriteLine|ReadLine|TryParse|forEach|includes|assign|freeze|splice|unshift|shift|flatMap)$"
|
||||||
|
},
|
||||||
|
"defect_class": "hallucination",
|
||||||
|
"inline_tier": "warning",
|
||||||
|
"filePath": "/home/alex/.npm-global/lib/node_modules/pi-lens/rules/tree-sitter-queries/python/python-cross-language-method.yml"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": "python-debugger",
|
||||||
|
"name": "Debugger Statement",
|
||||||
|
"severity": "warning",
|
||||||
|
"language": "python",
|
||||||
|
"message": "Debugger call '{{FUNC}}' — remove before committing",
|
||||||
|
"query": " (call\n function: (identifier) @FUNC\n (#eq? @FUNC \"breakpoint\"))\n\n (call\n function: (attribute\n object: (identifier) @MOD\n attribute: (identifier) @FUNC)\n (#eq? @MOD \"pdb\")\n (#match? @FUNC \"^(set_trace|post_mortem|pm|run|runcall)$\"))",
|
||||||
|
"metavars": [
|
||||||
|
"FUNC",
|
||||||
|
"MOD"
|
||||||
|
],
|
||||||
|
"defect_class": "safety",
|
||||||
|
"inline_tier": "blocking",
|
||||||
|
"filePath": "/home/alex/.npm-global/lib/node_modules/pi-lens/rules/tree-sitter-queries/python/python-debugger.yml"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": "python-empty-except",
|
||||||
|
"name": "Empty Except Block",
|
||||||
|
"severity": "warning",
|
||||||
|
"language": "python",
|
||||||
|
"message": "Except block only contains 'pass' — handle or re-raise the exception",
|
||||||
|
"query": " (try_statement\n (except_clause\n body: (block) @BODY))",
|
||||||
|
"metavars": [
|
||||||
|
"BODY"
|
||||||
|
],
|
||||||
|
"post_filter": "python_empty_except",
|
||||||
|
"defect_class": "silent-error",
|
||||||
|
"inline_tier": "blocking",
|
||||||
|
"filePath": "/home/alex/.npm-global/lib/node_modules/pi-lens/rules/tree-sitter-queries/python/python-empty-except.yml"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": "python-hallucinated-import",
|
||||||
|
"name": "Hallucinated Import",
|
||||||
|
"severity": "warning",
|
||||||
|
"language": "python",
|
||||||
|
"message": "Hallucinated import — '{NAME}' does not exist in '{MODULE}'",
|
||||||
|
"query": " (import_from_statement\n module_name: (dotted_name) @MODULE\n name: (dotted_name) @NAME)",
|
||||||
|
"metavars": [
|
||||||
|
"MODULE",
|
||||||
|
"NAME"
|
||||||
|
],
|
||||||
|
"post_filter": "match_captures",
|
||||||
|
"post_filter_params": {
|
||||||
|
"MODULE": "^(requests|flask|django|typing|collections|asyncio|json|unittest|pytest|urllib|sqlalchemy)$",
|
||||||
|
"NAME": "^(JSONResponse|HTMLResponse|RedirectResponse|StreamingResponse|Depends|Query|Path|Body|Header|Cookie|Form|File|UploadFile|FastAPI|APIRouter|HTTPException|BackgroundTasks|dataclass|fields|BaseModel|Field|validator|aiohttp|parse|stringify|fixture|TestCase|get|post|put|delete|Model|Session|Column|Integer|String)$"
|
||||||
|
},
|
||||||
|
"defect_class": "hallucination",
|
||||||
|
"inline_tier": "blocking",
|
||||||
|
"filePath": "/home/alex/.npm-global/lib/node_modules/pi-lens/rules/tree-sitter-queries/python/python-hallucinated-import.yml"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": "python-hardcoded-secrets",
|
||||||
|
"name": "Hardcoded Secret",
|
||||||
|
"severity": "warning",
|
||||||
|
"language": "python",
|
||||||
|
"message": "Hardcoded {{VARNAME}} — use environment variables or a secrets manager",
|
||||||
|
"query": " (assignment\n left: (identifier) @VARNAME\n right: (string) @VALUE)",
|
||||||
|
"metavars": [
|
||||||
|
"VARNAME",
|
||||||
|
"VALUE"
|
||||||
|
],
|
||||||
|
"post_filter": "check_secret_pattern",
|
||||||
|
"defect_class": "secrets",
|
||||||
|
"inline_tier": "blocking",
|
||||||
|
"filePath": "/home/alex/.npm-global/lib/node_modules/pi-lens/rules/tree-sitter-queries/python/python-hardcoded-secrets.yml"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": "python-insecure-deserialization",
|
||||||
|
"name": "Insecure Deserialization",
|
||||||
|
"severity": "error",
|
||||||
|
"language": "python",
|
||||||
|
"message": "Potential insecure deserialization sink — avoid unsafe loaders",
|
||||||
|
"query": " (call\n function: (attribute\n object: (identifier) @MOD\n attribute: (identifier) @FN)\n arguments: (argument_list (_) @DATA)\n (#match? @MOD \"^(pickle|yaml)$\")\n (#match? @FN \"^(load|loads|unsafe_load)$\"))",
|
||||||
|
"metavars": [
|
||||||
|
"MOD",
|
||||||
|
"FN",
|
||||||
|
"DATA"
|
||||||
|
],
|
||||||
|
"post_filter": "py_insecure_deserialization_sink",
|
||||||
|
"defect_class": "injection",
|
||||||
|
"inline_tier": "blocking",
|
||||||
|
"filePath": "/home/alex/.npm-global/lib/node_modules/pi-lens/rules/tree-sitter-queries/python/python-insecure-deserialization.yml"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": "python-insecure-random",
|
||||||
|
"name": "Insecure Randomness",
|
||||||
|
"severity": "warning",
|
||||||
|
"language": "python",
|
||||||
|
"message": "Insecure randomness source detected — use secrets or os.urandom for security-sensitive values",
|
||||||
|
"query": " (call\n function: (attribute\n object: (identifier) @MOD\n attribute: (identifier) @FN)\n arguments: (argument_list) @ARGS\n (#eq? @MOD \"random\")\n (#match? @FN \"^(random|randint|randrange|choice|choices)$\"))",
|
||||||
|
"metavars": [
|
||||||
|
"MOD",
|
||||||
|
"FN",
|
||||||
|
"ARGS"
|
||||||
|
],
|
||||||
|
"defect_class": "injection",
|
||||||
|
"inline_tier": "warning",
|
||||||
|
"filePath": "/home/alex/.npm-global/lib/node_modules/pi-lens/rules/tree-sitter-queries/python/python-insecure-random.yml"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": "python-mutable-class-attr",
|
||||||
|
"name": "Mutable Class Attribute",
|
||||||
|
"severity": "warning",
|
||||||
|
"language": "python",
|
||||||
|
"message": "Class attribute '{{VARNAME}}' is mutable — shared across all instances",
|
||||||
|
"query": " (class_definition\n body: (block\n (expression_statement\n (assignment\n left: (identifier) @VARNAME\n right: [\n (list) @VALUE\n (dictionary) @VALUE\n (set) @VALUE\n ]))))",
|
||||||
|
"metavars": [
|
||||||
|
"VARNAME",
|
||||||
|
"VALUE"
|
||||||
|
],
|
||||||
|
"post_filter": "not_in_function",
|
||||||
|
"defect_class": "correctness",
|
||||||
|
"inline_tier": "blocking",
|
||||||
|
"filePath": "/home/alex/.npm-global/lib/node_modules/pi-lens/rules/tree-sitter-queries/python/python-mutable-class-attr.yml"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": "python-path-traversal",
|
||||||
|
"name": "Path Traversal Risk",
|
||||||
|
"severity": "warning",
|
||||||
|
"language": "python",
|
||||||
|
"message": "Potential path traversal sink — sanitize and constrain file paths",
|
||||||
|
"query": " [\n (call\n function: (identifier) @FN\n arguments: (argument_list\n [(identifier) (binary_operator) (call)] @PATH))\n (call\n function: (attribute\n object: (identifier) @MOD\n attribute: (identifier) @FN)\n arguments: (argument_list\n [(identifier) (binary_operator) (call)] @PATH))\n ]\n (#match? @FN \"^(open|read_text|read_bytes|write_text|write_bytes|remove|unlink|rmdir)$\")",
|
||||||
|
"metavars": [
|
||||||
|
"MOD",
|
||||||
|
"FN",
|
||||||
|
"PATH"
|
||||||
|
],
|
||||||
|
"post_filter": "py_path_traversal_sink",
|
||||||
|
"defect_class": "injection",
|
||||||
|
"inline_tier": "warning",
|
||||||
|
"filePath": "/home/alex/.npm-global/lib/node_modules/pi-lens/rules/tree-sitter-queries/python/python-path-traversal.yml"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": "python-print-statement",
|
||||||
|
"name": "Print Statement in Production",
|
||||||
|
"severity": "warning",
|
||||||
|
"language": "python",
|
||||||
|
"message": "print() — remove debug output before committing",
|
||||||
|
"query": " (call\n function: (identifier) @FUNC\n (#eq? @FUNC \"print\")\n arguments: (argument_list) @ARGS)",
|
||||||
|
"metavars": [
|
||||||
|
"FUNC",
|
||||||
|
"ARGS"
|
||||||
|
],
|
||||||
|
"post_filter": "not_in_test_block",
|
||||||
|
"defect_class": "safety",
|
||||||
|
"inline_tier": "warning",
|
||||||
|
"filePath": "/home/alex/.npm-global/lib/node_modules/pi-lens/rules/tree-sitter-queries/python/python-print-statement.yml"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": "python-raise-string",
|
||||||
|
"name": "Raise String Instead of Exception",
|
||||||
|
"severity": "warning",
|
||||||
|
"language": "python",
|
||||||
|
"message": "raise with string literal — Python 3 requires exception instances",
|
||||||
|
"query": " (raise_statement\n (string) @VALUE)",
|
||||||
|
"metavars": [
|
||||||
|
"VALUE"
|
||||||
|
],
|
||||||
|
"defect_class": "correctness",
|
||||||
|
"inline_tier": "blocking",
|
||||||
|
"filePath": "/home/alex/.npm-global/lib/node_modules/pi-lens/rules/tree-sitter-queries/python/python-raise-string.yml"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": "python-sleep-in-test",
|
||||||
|
"name": "time.sleep in Test",
|
||||||
|
"severity": "warning",
|
||||||
|
"language": "python",
|
||||||
|
"message": "time.sleep() in test — use synchronisation primitives or polling helpers instead of fixed sleeps",
|
||||||
|
"query": " (call\n function: (attribute\n object: (identifier) @MOD\n attribute: (identifier) @FN)\n (#eq? @MOD \"time\")\n (#eq? @FN \"sleep\")) @CALL",
|
||||||
|
"metavars": [
|
||||||
|
"MOD",
|
||||||
|
"FN",
|
||||||
|
"CALL"
|
||||||
|
],
|
||||||
|
"defect_class": "async-misuse",
|
||||||
|
"inline_tier": "warning",
|
||||||
|
"filePath": "/home/alex/.npm-global/lib/node_modules/pi-lens/rules/tree-sitter-queries/python/python-sleep-in-test.yml"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": "python-sql-injection",
|
||||||
|
"name": "SQL Injection Risk",
|
||||||
|
"severity": "error",
|
||||||
|
"language": "python",
|
||||||
|
"message": "Potential SQL injection sink — use parameterized queries",
|
||||||
|
"query": " (call\n function: (attribute\n object: (_) @OBJ\n attribute: (identifier) @FN)\n arguments: (argument_list\n [(binary_operator) (identifier) (call)] @SQL\n (_)*))",
|
||||||
|
"metavars": [
|
||||||
|
"OBJ",
|
||||||
|
"FN",
|
||||||
|
"SQL"
|
||||||
|
],
|
||||||
|
"post_filter": "py_sql_injection_sink",
|
||||||
|
"defect_class": "injection",
|
||||||
|
"inline_tier": "blocking",
|
||||||
|
"filePath": "/home/alex/.npm-global/lib/node_modules/pi-lens/rules/tree-sitter-queries/python/python-sql-injection.yml"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": "python-ssrf",
|
||||||
|
"name": "SSRF Risk",
|
||||||
|
"severity": "warning",
|
||||||
|
"language": "python",
|
||||||
|
"message": "Potential SSRF sink — validate/allowlist outbound URLs",
|
||||||
|
"query": " (call\n function: (attribute\n object: (identifier) @MOD\n attribute: (identifier) @FN)\n arguments: (argument_list\n [(identifier) (subscript) (call)] @URL)\n (#eq? @MOD \"requests\")\n (#match? @FN \"^(get|post|put|patch|delete|request|head|options)$\"))",
|
||||||
|
"metavars": [
|
||||||
|
"MOD",
|
||||||
|
"FN",
|
||||||
|
"URL"
|
||||||
|
],
|
||||||
|
"post_filter": "py_ssrf_sink",
|
||||||
|
"defect_class": "injection",
|
||||||
|
"inline_tier": "warning",
|
||||||
|
"filePath": "/home/alex/.npm-global/lib/node_modules/pi-lens/rules/tree-sitter-queries/python/python-ssrf.yml"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": "python-subprocess-shell",
|
||||||
|
"name": "subprocess with shell=True",
|
||||||
|
"severity": "warning",
|
||||||
|
"language": "python",
|
||||||
|
"message": "subprocess called with shell=True — command injection risk if any argument is user-controlled",
|
||||||
|
"query": " (call\n function: (attribute\n object: (identifier) @MOD\n attribute: (identifier) @FN)\n arguments: (argument_list\n (keyword_argument\n name: (identifier) @KW\n value: (true) @VAL))\n (#eq? @MOD \"subprocess\")\n (#match? @FN \"^(run|Popen|call|check_output|check_call)$\")\n (#eq? @KW \"shell\"))",
|
||||||
|
"metavars": [
|
||||||
|
"MOD",
|
||||||
|
"FN",
|
||||||
|
"KW"
|
||||||
|
],
|
||||||
|
"defect_class": "injection",
|
||||||
|
"inline_tier": "warning",
|
||||||
|
"filePath": "/home/alex/.npm-global/lib/node_modules/pi-lens/rules/tree-sitter-queries/python/python-subprocess-shell.yml"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": "python-thread-global-write",
|
||||||
|
"name": "Threaded Shared State Risk",
|
||||||
|
"severity": "warning",
|
||||||
|
"language": "python",
|
||||||
|
"message": "Thread creation detected — ensure shared state mutations are synchronized",
|
||||||
|
"query": " (call\n function: (attribute\n object: (identifier) @MOD\n attribute: (identifier) @FN)\n arguments: (argument_list) @ARGS)\n (#eq? @MOD \"threading\")\n (#eq? @FN \"Thread\")",
|
||||||
|
"metavars": [
|
||||||
|
"MOD",
|
||||||
|
"FN",
|
||||||
|
"ARGS"
|
||||||
|
],
|
||||||
|
"defect_class": "async-misuse",
|
||||||
|
"inline_tier": "warning",
|
||||||
|
"filePath": "/home/alex/.npm-global/lib/node_modules/pi-lens/rules/tree-sitter-queries/python/python-thread-global-write.yml"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": "python-unsafe-regex",
|
||||||
|
"name": "Unsafe Dynamic Regex",
|
||||||
|
"severity": "warning",
|
||||||
|
"language": "python",
|
||||||
|
"message": "re.{{FUNC}}() with variable pattern — ReDoS risk if pattern is user-controlled",
|
||||||
|
"query": " (call\n function: (attribute\n object: (identifier) @MOD\n attribute: (identifier) @FUNC)\n arguments: (argument_list\n (identifier) @PATTERN)\n (#eq? @MOD \"re\")\n (#match? @FUNC \"^(compile|match|search|fullmatch|findall|finditer|sub|subn|split)$\"))",
|
||||||
|
"metavars": [
|
||||||
|
"MOD",
|
||||||
|
"FUNC",
|
||||||
|
"PATTERN"
|
||||||
|
],
|
||||||
|
"defect_class": "injection",
|
||||||
|
"inline_tier": "blocking",
|
||||||
|
"filePath": "/home/alex/.npm-global/lib/node_modules/pi-lens/rules/tree-sitter-queries/python/python-unsafe-regex.yml"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": "python-weak-hash",
|
||||||
|
"name": "Weak Hash Primitive",
|
||||||
|
"severity": "error",
|
||||||
|
"language": "python",
|
||||||
|
"message": "Weak hash primitive detected (MD5/SHA1) — use SHA-256+ for security-sensitive contexts",
|
||||||
|
"query": " (call\n function: (attribute\n object: (identifier) @MOD\n attribute: (identifier) @FN)\n arguments: (argument_list) @ARGS\n (#eq? @MOD \"hashlib\")\n (#match? @FN \"^(md5|sha1)$\"))",
|
||||||
|
"metavars": [
|
||||||
|
"MOD",
|
||||||
|
"FN",
|
||||||
|
"ARGS"
|
||||||
|
],
|
||||||
|
"defect_class": "injection",
|
||||||
|
"inline_tier": "blocking",
|
||||||
|
"filePath": "/home/alex/.npm-global/lib/node_modules/pi-lens/rules/tree-sitter-queries/python/python-weak-hash.yml"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": "return-in-generator",
|
||||||
|
"name": "Return with Value in Generator",
|
||||||
|
"severity": "error",
|
||||||
|
"language": "python",
|
||||||
|
"message": "'return' with a value should not be used in a generator function",
|
||||||
|
"query": " (function_definition\n body: (block\n (return_statement\n (_) @RETURN_VAL) @RETURN)) @FUNCTION",
|
||||||
|
"metavars": [
|
||||||
|
"FUNCTION",
|
||||||
|
"RETURN",
|
||||||
|
"RETURN_VAL"
|
||||||
|
],
|
||||||
|
"post_filter": "is_generator_with_valued_return",
|
||||||
|
"defect_class": "correctness",
|
||||||
|
"inline_tier": "blocking",
|
||||||
|
"filePath": "/home/alex/.npm-global/lib/node_modules/pi-lens/rules/tree-sitter-queries/python/return-in-generator.yml"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": "return-in-init",
|
||||||
|
"name": "Return Value in __init__",
|
||||||
|
"severity": "error",
|
||||||
|
"language": "python",
|
||||||
|
"message": "__init__ should not return a value — it must always return None",
|
||||||
|
"query": " (function_definition\n name: (identifier) @NAME (#eq? @NAME \"__init__\")\n body: (block\n (return_statement\n (_) @RETURN_VAL) @RETURN))",
|
||||||
|
"metavars": [
|
||||||
|
"NAME",
|
||||||
|
"RETURN",
|
||||||
|
"RETURN_VAL"
|
||||||
|
],
|
||||||
|
"post_filter": "has_return_value",
|
||||||
|
"defect_class": "correctness",
|
||||||
|
"inline_tier": "blocking",
|
||||||
|
"filePath": "/home/alex/.npm-global/lib/node_modules/pi-lens/rules/tree-sitter-queries/python/return-in-init.yml"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": "send-file-mimetype",
|
||||||
|
"name": "send_file Should Specify Mimetype or Download Name",
|
||||||
|
"severity": "error",
|
||||||
|
"language": "python",
|
||||||
|
"message": "send_file should specify 'mimetype' or 'download_name' when used with file-like objects",
|
||||||
|
"query": " (call\n function: (identifier) @FUNC (#eq? @FUNC \"send_file\")\n arguments: (argument_list\n (_) @FIRST_ARG\n (keyword_argument)? @KW))",
|
||||||
|
"metavars": [
|
||||||
|
"FUNC",
|
||||||
|
"FIRST_ARG",
|
||||||
|
"KW"
|
||||||
|
],
|
||||||
|
"post_filter": "missing_mimetype_and_download_name",
|
||||||
|
"defect_class": "correctness",
|
||||||
|
"inline_tier": "blocking",
|
||||||
|
"filePath": "/home/alex/.npm-global/lib/node_modules/pi-lens/rules/tree-sitter-queries/python/send-file-mimetype.yml"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": "unreachable-except",
|
||||||
|
"name": "Unreachable Except Clause",
|
||||||
|
"severity": "warning",
|
||||||
|
"language": "python",
|
||||||
|
"message": "Unreachable except clause — earlier except catches all",
|
||||||
|
"query": " (try_statement\n (except_clause\n \"except\") @GENERAL\n (except_clause\n \"except\"\n (identifier) @SPECIFIC))",
|
||||||
|
"metavars": [
|
||||||
|
"GENERAL",
|
||||||
|
"SPECIFIC"
|
||||||
|
],
|
||||||
|
"defect_class": "correctness",
|
||||||
|
"inline_tier": "blocking",
|
||||||
|
"filePath": "/home/alex/.npm-global/lib/node_modules/pi-lens/rules/tree-sitter-queries/python/unreachable-except.yml"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": "wildcard-import",
|
||||||
|
"name": "Wildcard Import",
|
||||||
|
"severity": "warning",
|
||||||
|
"language": "python",
|
||||||
|
"message": "Wildcard import — pollutes namespace, hard to track origin",
|
||||||
|
"query": " (import_from_statement\n module_name: (dotted_name) @MODULE\n (wildcard_import) @WILDCARD)",
|
||||||
|
"metavars": [
|
||||||
|
"MODULE",
|
||||||
|
"WILDCARD"
|
||||||
|
],
|
||||||
|
"defect_class": "safety",
|
||||||
|
"inline_tier": "warning",
|
||||||
|
"filePath": "/home/alex/.npm-global/lib/node_modules/pi-lens/rules/tree-sitter-queries/python/wildcard-import.yml"
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"id": "yield-return-outside-function",
|
||||||
|
"name": "Yield/Return Outside Function",
|
||||||
|
"severity": "error",
|
||||||
|
"language": "python",
|
||||||
|
"message": "{{STATEMENT}} used outside function — syntax error",
|
||||||
|
"query": " (module\n (expression_statement\n (yield) @STATEMENT))\n (module\n (expression_statement\n (yield_expression) @STATEMENT))\n (module\n (return_statement) @STATEMENT)",
|
||||||
|
"metavars": [
|
||||||
|
"STATEMENT"
|
||||||
|
],
|
||||||
|
"defect_class": "correctness",
|
||||||
|
"inline_tier": "blocking",
|
||||||
|
"filePath": "/home/alex/.npm-global/lib/node_modules/pi-lens/rules/tree-sitter-queries/python/yield-return-outside-function.yml"
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
@@ -7,8 +7,6 @@ Create Date: 2026-05-22 21:50:00.000000
|
|||||||
"""
|
"""
|
||||||
from typing import Sequence, Union
|
from typing import Sequence, Union
|
||||||
|
|
||||||
from alembic import op
|
|
||||||
import sqlalchemy as sa
|
|
||||||
|
|
||||||
# revision identifiers, used by Alembic.
|
# revision identifiers, used by Alembic.
|
||||||
revision: str = "0014_merge_heads"
|
revision: str = "0014_merge_heads"
|
||||||
|
|||||||
@@ -10,7 +10,6 @@ from typing import Sequence, Union
|
|||||||
from alembic import op
|
from alembic import op
|
||||||
import sqlalchemy as sa
|
import sqlalchemy as sa
|
||||||
from sqlalchemy.dialects import postgresql
|
from sqlalchemy.dialects import postgresql
|
||||||
from sqlalchemy import inspect
|
|
||||||
|
|
||||||
# revision identifiers, used by Alembic.
|
# revision identifiers, used by Alembic.
|
||||||
revision: str = "0015_single_interface"
|
revision: str = "0015_single_interface"
|
||||||
|
|||||||
@@ -0,0 +1,32 @@
|
|||||||
|
"""add_ssh_key_id_to_config_profiles
|
||||||
|
|
||||||
|
Revision ID: 069d3da4dc9b
|
||||||
|
Revises: 2026_05_29_add_notifications_table
|
||||||
|
Create Date: 2026-05-29 12:30:16.580532
|
||||||
|
"""
|
||||||
|
|
||||||
|
from alembic import op
|
||||||
|
import sqlalchemy as sa
|
||||||
|
|
||||||
|
|
||||||
|
# revision identifiers, used by Alembic.
|
||||||
|
revision = "069d3da4dc9b"
|
||||||
|
down_revision = "2026_05_29_add_notifications_table"
|
||||||
|
branch_labels = None
|
||||||
|
depends_on = None
|
||||||
|
|
||||||
|
|
||||||
|
def upgrade() -> None:
|
||||||
|
op.add_column(
|
||||||
|
"config_profiles",
|
||||||
|
sa.Column(
|
||||||
|
"ssh_key_id",
|
||||||
|
sa.Uuid(),
|
||||||
|
sa.ForeignKey("ssh_keys.id", ondelete="SET NULL"),
|
||||||
|
nullable=True,
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def downgrade() -> None:
|
||||||
|
op.drop_column("config_profiles", "ssh_key_id")
|
||||||
@@ -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,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,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,122 @@
|
|||||||
|
"""add monitoring tables
|
||||||
|
|
||||||
|
Revision ID: 2026_05_28_add_monitoring_tables
|
||||||
|
Revises: 2026_05_28_drop_tool_configs_and_config_folders
|
||||||
|
Create Date: 2026-05-28
|
||||||
|
|
||||||
|
"""
|
||||||
|
|
||||||
|
from collections.abc import Sequence
|
||||||
|
|
||||||
|
from alembic import op
|
||||||
|
import sqlalchemy as sa
|
||||||
|
|
||||||
|
# revision identifiers, used by Alembic.
|
||||||
|
revision: str = "2026_05_28_add_monitoring_tables"
|
||||||
|
down_revision: str | None = "2026_05_28_drop_tool_configs_and_config_folders"
|
||||||
|
branch_labels: str | Sequence[str] | None = None
|
||||||
|
depends_on: str | Sequence[str] | None = None
|
||||||
|
|
||||||
|
|
||||||
|
def upgrade() -> None:
|
||||||
|
op.create_table(
|
||||||
|
"instance_events",
|
||||||
|
sa.Column("id", sa.Uuid(), nullable=False),
|
||||||
|
sa.Column(
|
||||||
|
"instance_id",
|
||||||
|
sa.Uuid(),
|
||||||
|
nullable=False,
|
||||||
|
),
|
||||||
|
sa.Column("event_type", sa.String(length=50), nullable=False),
|
||||||
|
sa.Column("status", sa.String(length=50), nullable=True),
|
||||||
|
sa.Column("message", sa.Text(), nullable=True),
|
||||||
|
sa.Column("created_by", sa.Uuid(), nullable=True),
|
||||||
|
sa.Column(
|
||||||
|
"metadata",
|
||||||
|
sa.JSON(),
|
||||||
|
nullable=False,
|
||||||
|
server_default="{}",
|
||||||
|
),
|
||||||
|
sa.Column(
|
||||||
|
"created_at",
|
||||||
|
sa.DateTime(timezone=True),
|
||||||
|
server_default=sa.func.now(),
|
||||||
|
nullable=False,
|
||||||
|
),
|
||||||
|
sa.ForeignKeyConstraint(
|
||||||
|
["instance_id"],
|
||||||
|
["tool_instances.id"],
|
||||||
|
ondelete="CASCADE",
|
||||||
|
),
|
||||||
|
sa.ForeignKeyConstraint(
|
||||||
|
["created_by"],
|
||||||
|
["users.id"],
|
||||||
|
ondelete="SET NULL",
|
||||||
|
),
|
||||||
|
sa.PrimaryKeyConstraint("id"),
|
||||||
|
)
|
||||||
|
op.create_index(
|
||||||
|
"idx_instance_events_instance_id",
|
||||||
|
"instance_events",
|
||||||
|
["instance_id"],
|
||||||
|
)
|
||||||
|
op.create_index(
|
||||||
|
"idx_instance_events_created_at",
|
||||||
|
"instance_events",
|
||||||
|
["created_at"],
|
||||||
|
postgresql_using="btree",
|
||||||
|
)
|
||||||
|
op.create_index(
|
||||||
|
"idx_instance_events_event_type",
|
||||||
|
"instance_events",
|
||||||
|
["event_type"],
|
||||||
|
)
|
||||||
|
|
||||||
|
op.create_table(
|
||||||
|
"health_checks",
|
||||||
|
sa.Column("id", sa.Uuid(), nullable=False),
|
||||||
|
sa.Column(
|
||||||
|
"instance_id",
|
||||||
|
sa.Uuid(),
|
||||||
|
nullable=False,
|
||||||
|
),
|
||||||
|
sa.Column("container_status", sa.String(length=50), nullable=True),
|
||||||
|
sa.Column("container_healthy", sa.Boolean(), nullable=True),
|
||||||
|
sa.Column("tunnel_healthy", sa.Boolean(), nullable=True),
|
||||||
|
sa.Column("exit_code", sa.Integer(), nullable=True),
|
||||||
|
sa.Column("probe_status", sa.String(length=50), nullable=True),
|
||||||
|
sa.Column("probe_output", sa.Text(), nullable=True),
|
||||||
|
sa.Column(
|
||||||
|
"checked_at",
|
||||||
|
sa.DateTime(timezone=True),
|
||||||
|
server_default=sa.func.now(),
|
||||||
|
nullable=False,
|
||||||
|
),
|
||||||
|
sa.ForeignKeyConstraint(
|
||||||
|
["instance_id"],
|
||||||
|
["tool_instances.id"],
|
||||||
|
ondelete="CASCADE",
|
||||||
|
),
|
||||||
|
sa.PrimaryKeyConstraint("id"),
|
||||||
|
)
|
||||||
|
op.create_index(
|
||||||
|
"idx_health_checks_instance_id",
|
||||||
|
"health_checks",
|
||||||
|
["instance_id"],
|
||||||
|
)
|
||||||
|
op.create_index(
|
||||||
|
"idx_health_checks_checked_at",
|
||||||
|
"health_checks",
|
||||||
|
["checked_at"],
|
||||||
|
postgresql_using="btree",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def downgrade() -> None:
|
||||||
|
op.drop_index("idx_health_checks_checked_at", table_name="health_checks")
|
||||||
|
op.drop_index("idx_health_checks_instance_id", table_name="health_checks")
|
||||||
|
op.drop_table("health_checks")
|
||||||
|
op.drop_index("idx_instance_events_event_type", table_name="instance_events")
|
||||||
|
op.drop_index("idx_instance_events_created_at", table_name="instance_events")
|
||||||
|
op.drop_index("idx_instance_events_instance_id", table_name="instance_events")
|
||||||
|
op.drop_table("instance_events")
|
||||||
@@ -0,0 +1,61 @@
|
|||||||
|
"""add terminal_sessions table
|
||||||
|
|
||||||
|
Revision ID: 2026_05_28_add_terminal_sessions
|
||||||
|
Revises: 20260527_160017_add_pi_agent
|
||||||
|
Create Date: 2026-05-28
|
||||||
|
|
||||||
|
"""
|
||||||
|
|
||||||
|
from collections.abc import Sequence
|
||||||
|
|
||||||
|
from alembic import op
|
||||||
|
import sqlalchemy as sa
|
||||||
|
|
||||||
|
# revision identifiers, used by Alembic.
|
||||||
|
revision: str = "2026_05_28_add_terminal_sessions"
|
||||||
|
down_revision: str | None = "2026_05_28_add_tool_definition_manifests"
|
||||||
|
branch_labels: str | Sequence[str] | None = None
|
||||||
|
depends_on: str | Sequence[str] | None = None
|
||||||
|
|
||||||
|
|
||||||
|
def upgrade() -> None:
|
||||||
|
op.create_table(
|
||||||
|
"terminal_sessions",
|
||||||
|
sa.Column("id", sa.UUID(), nullable=False),
|
||||||
|
sa.Column("instance_id", sa.UUID(), nullable=False),
|
||||||
|
sa.Column("name", sa.String(length=255), nullable=True),
|
||||||
|
sa.Column("status", sa.String(length=50), nullable=False),
|
||||||
|
sa.Column("last_activity_at", sa.DateTime(timezone=True), nullable=True),
|
||||||
|
sa.Column("closed_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()"),
|
||||||
|
onupdate=sa.text("now()"),
|
||||||
|
nullable=False,
|
||||||
|
),
|
||||||
|
sa.ForeignKeyConstraint(
|
||||||
|
["instance_id"], ["tool_instances.id"], ondelete="CASCADE"
|
||||||
|
),
|
||||||
|
sa.PrimaryKeyConstraint("id"),
|
||||||
|
)
|
||||||
|
op.create_index(
|
||||||
|
op.f("ix_terminal_sessions_instance_id"),
|
||||||
|
"terminal_sessions",
|
||||||
|
["instance_id"],
|
||||||
|
unique=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def downgrade() -> None:
|
||||||
|
op.drop_index(
|
||||||
|
op.f("ix_terminal_sessions_instance_id"),
|
||||||
|
table_name="terminal_sessions",
|
||||||
|
)
|
||||||
|
op.drop_table("terminal_sessions")
|
||||||
@@ -0,0 +1,373 @@
|
|||||||
|
"""add tool definition manifests
|
||||||
|
|
||||||
|
Revision ID: 2026_05_28_add_tool_definition_manifests
|
||||||
|
Revises: 20260527_160017_add_pi_agent
|
||||||
|
Create Date: 2026-05-28T11:00:00
|
||||||
|
|
||||||
|
"""
|
||||||
|
|
||||||
|
import json
|
||||||
|
import uuid
|
||||||
|
from typing import Sequence, Union
|
||||||
|
|
||||||
|
from alembic import op
|
||||||
|
import sqlalchemy as sa
|
||||||
|
|
||||||
|
# revision identifiers, used by Alembic.
|
||||||
|
revision: str = "2026_05_28_add_tool_definition_manifests"
|
||||||
|
down_revision: Union[str, None] = "20260527_160017_add_pi_agent"
|
||||||
|
branch_labels: Union[str, Sequence[str], None] = None
|
||||||
|
depends_on: Union[str, Sequence[str], None] = None
|
||||||
|
|
||||||
|
BASE_UBUNTU_ID = uuid.UUID("a1b2c3d4-e5f6-7890-abcd-ef1234567890")
|
||||||
|
PI_AGENT_MANIFEST_ID = uuid.UUID("d07b8376-2151-4119-8c1d-27f792aae9a3")
|
||||||
|
|
||||||
|
|
||||||
|
def upgrade() -> None:
|
||||||
|
conn = op.get_bind()
|
||||||
|
|
||||||
|
# ── Create tool_definition_manifests table ───────────────────────
|
||||||
|
op.create_table(
|
||||||
|
"tool_definition_manifests",
|
||||||
|
sa.Column("id", sa.UUID(), nullable=False),
|
||||||
|
sa.Column("name", sa.String(64), nullable=False),
|
||||||
|
sa.Column("display_name", sa.String(128), nullable=False),
|
||||||
|
sa.Column("description", sa.Text(), nullable=True),
|
||||||
|
sa.Column("category", sa.String(64), nullable=True),
|
||||||
|
sa.Column("interface_type", sa.String(16), nullable=False),
|
||||||
|
sa.Column("base_image", sa.String(256), nullable=True),
|
||||||
|
sa.Column("base_definition_id", sa.UUID(), nullable=True),
|
||||||
|
sa.Column(
|
||||||
|
"base_version", sa.String(32), nullable=False, server_default="latest"
|
||||||
|
),
|
||||||
|
sa.Column("manifest", sa.JSON(), nullable=False),
|
||||||
|
sa.Column("dockerfile_cache", sa.Text(), nullable=True),
|
||||||
|
sa.Column("compose_cache", sa.Text(), nullable=True),
|
||||||
|
sa.Column("version", sa.String(32), nullable=False, server_default="v1"),
|
||||||
|
sa.Column("is_base", sa.Boolean(), nullable=False, server_default="false"),
|
||||||
|
sa.Column("created_by_id", sa.UUID(), nullable=True),
|
||||||
|
sa.Column(
|
||||||
|
"created_at", sa.TIMESTAMP(timezone=True), server_default=sa.func.now()
|
||||||
|
),
|
||||||
|
sa.Column(
|
||||||
|
"updated_at", sa.TIMESTAMP(timezone=True), server_default=sa.func.now()
|
||||||
|
),
|
||||||
|
sa.PrimaryKeyConstraint("id"),
|
||||||
|
sa.UniqueConstraint("name"),
|
||||||
|
sa.ForeignKeyConstraint(
|
||||||
|
["base_definition_id"], ["tool_definition_manifests.id"]
|
||||||
|
),
|
||||||
|
sa.ForeignKeyConstraint(["created_by_id"], ["users.id"]),
|
||||||
|
sa.CheckConstraint(
|
||||||
|
"(base_image IS NOT NULL) OR (base_definition_id IS NOT NULL)",
|
||||||
|
name="ck_tool_definition_manifests_base_required",
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
# ── Add columns to tool_types ────────────────────────────────────
|
||||||
|
# Check if manifest_id exists before adding
|
||||||
|
conn = op.get_bind()
|
||||||
|
result = conn.execute(
|
||||||
|
sa.text("""
|
||||||
|
SELECT column_name FROM information_schema.columns
|
||||||
|
WHERE table_name = 'tool_types' AND column_name = 'manifest_id'
|
||||||
|
""")
|
||||||
|
)
|
||||||
|
if not result.fetchone():
|
||||||
|
op.add_column("tool_types", sa.Column("manifest_id", sa.UUID(), nullable=True))
|
||||||
|
op.create_foreign_key(
|
||||||
|
"fk_tool_types_manifest_id",
|
||||||
|
"tool_types",
|
||||||
|
"tool_definition_manifests",
|
||||||
|
["manifest_id"],
|
||||||
|
["id"],
|
||||||
|
)
|
||||||
|
|
||||||
|
# Update definition_type to allow 'legacy' and 'manifest'
|
||||||
|
result = conn.execute(
|
||||||
|
sa.text("""
|
||||||
|
SELECT constraint_name FROM information_schema.check_constraints
|
||||||
|
WHERE constraint_name = 'chk_definition_type'
|
||||||
|
""")
|
||||||
|
)
|
||||||
|
if result.fetchone():
|
||||||
|
op.drop_constraint("chk_definition_type", "tool_types", type_="check")
|
||||||
|
|
||||||
|
op.execute("ALTER TABLE tool_types ALTER COLUMN definition_type TYPE VARCHAR(16)")
|
||||||
|
op.execute(
|
||||||
|
"ALTER TABLE tool_types ALTER COLUMN definition_type SET DEFAULT 'legacy'"
|
||||||
|
)
|
||||||
|
|
||||||
|
# ── Add columns to tool_instances ────────────────────────────────
|
||||||
|
result = conn.execute(
|
||||||
|
sa.text("""
|
||||||
|
SELECT column_name FROM information_schema.columns
|
||||||
|
WHERE table_name = 'tool_instances' AND column_name = 'manifest_compiled_at'
|
||||||
|
""")
|
||||||
|
)
|
||||||
|
if not result.fetchone():
|
||||||
|
op.add_column(
|
||||||
|
"tool_instances",
|
||||||
|
sa.Column(
|
||||||
|
"manifest_compiled_at", sa.TIMESTAMP(timezone=True), nullable=True
|
||||||
|
),
|
||||||
|
)
|
||||||
|
|
||||||
|
result = conn.execute(
|
||||||
|
sa.text("""
|
||||||
|
SELECT column_name FROM information_schema.columns
|
||||||
|
WHERE table_name = 'tool_instances' AND column_name = 'image_tag'
|
||||||
|
""")
|
||||||
|
)
|
||||||
|
if not result.fetchone():
|
||||||
|
op.add_column(
|
||||||
|
"tool_instances",
|
||||||
|
sa.Column("image_tag", sa.String(256), nullable=True),
|
||||||
|
)
|
||||||
|
|
||||||
|
# ── Data migration: create base definition + pi-agent manifest ───
|
||||||
|
conn.execute(
|
||||||
|
sa.text(
|
||||||
|
"""
|
||||||
|
INSERT INTO tool_definition_manifests
|
||||||
|
(id, name, display_name, description, interface_type, base_image,
|
||||||
|
manifest, is_base, version, created_at, updated_at)
|
||||||
|
VALUES
|
||||||
|
(:base_id, 'ubuntu-24.04-dev', 'Ubuntu 24.04 Dev Base',
|
||||||
|
'Base development environment with build tools', 'terminal',
|
||||||
|
'ubuntu:24.04', :base_manifest, true, 'v1', now(), now())
|
||||||
|
"""
|
||||||
|
),
|
||||||
|
{
|
||||||
|
"base_id": BASE_UBUNTU_ID,
|
||||||
|
"base_manifest": json.dumps(
|
||||||
|
{
|
||||||
|
"name": "ubuntu-24.04-dev",
|
||||||
|
"display_name": "Ubuntu 24.04 Dev Base",
|
||||||
|
"interface_type": "terminal",
|
||||||
|
"base_image": "ubuntu:24.04",
|
||||||
|
"packages": {
|
||||||
|
"apt": [
|
||||||
|
"curl",
|
||||||
|
"wget",
|
||||||
|
"git",
|
||||||
|
"build-essential",
|
||||||
|
"ca-certificates",
|
||||||
|
"python3",
|
||||||
|
"python3-pip",
|
||||||
|
]
|
||||||
|
},
|
||||||
|
"user": {
|
||||||
|
"name": "user",
|
||||||
|
"uid": 1000,
|
||||||
|
"gid": 1000,
|
||||||
|
"create_home": True,
|
||||||
|
"shell": "/bin/bash",
|
||||||
|
},
|
||||||
|
"env": {"DEBIAN_FRONTEND": "noninteractive"},
|
||||||
|
}
|
||||||
|
),
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
conn.execute(
|
||||||
|
sa.text(
|
||||||
|
"""
|
||||||
|
INSERT INTO tool_definition_manifests
|
||||||
|
(id, name, display_name, description, category, interface_type,
|
||||||
|
base_definition_id, base_version, manifest, version, created_at, updated_at)
|
||||||
|
VALUES
|
||||||
|
(:manifest_id, 'pi-agent', 'Pi Agent',
|
||||||
|
'Terminal-based coding harness with nvim, ranger, tmux',
|
||||||
|
'development', 'terminal', :base_id, 'v1', :manifest, 'v1',
|
||||||
|
now(), now())
|
||||||
|
"""
|
||||||
|
),
|
||||||
|
{
|
||||||
|
"manifest_id": PI_AGENT_MANIFEST_ID,
|
||||||
|
"base_id": BASE_UBUNTU_ID,
|
||||||
|
"manifest": json.dumps(
|
||||||
|
{
|
||||||
|
"name": "pi-agent",
|
||||||
|
"display_name": "Pi Agent",
|
||||||
|
"description": "Terminal-based coding harness",
|
||||||
|
"category": "development",
|
||||||
|
"interface_type": "terminal",
|
||||||
|
"base_definition_id": str(BASE_UBUNTU_ID),
|
||||||
|
"base_version": "v1",
|
||||||
|
"packages": {
|
||||||
|
"apt": [
|
||||||
|
"neovim",
|
||||||
|
"ranger",
|
||||||
|
"tmux",
|
||||||
|
"htop",
|
||||||
|
"tree",
|
||||||
|
"jq",
|
||||||
|
],
|
||||||
|
"node": {"version": "20"},
|
||||||
|
"npm_global": ["@earendil-works/pi-coding-agent"],
|
||||||
|
},
|
||||||
|
"user": {
|
||||||
|
"name": "user",
|
||||||
|
"uid": 1001,
|
||||||
|
"gid": 1001,
|
||||||
|
"create_home": True,
|
||||||
|
"shell": "/bin/bash",
|
||||||
|
},
|
||||||
|
"env": {"DEBIAN_FRONTEND": "noninteractive"},
|
||||||
|
"scripts": {
|
||||||
|
"build": [
|
||||||
|
"git config --global init.defaultBranch main && git config --global user.email 'dev@headquarter.local' && git config --global user.name 'Developer'",
|
||||||
|
"mkdir -p /home/user/.config/ranger && echo 'set preview_files true' > /home/user/.config/ranger/rc.conf",
|
||||||
|
],
|
||||||
|
"startup": [
|
||||||
|
"if [ -d /workspace ]; then sudo chown -R user:user /workspace 2>/dev/null || true; fi",
|
||||||
|
],
|
||||||
|
},
|
||||||
|
"mounts": [
|
||||||
|
{
|
||||||
|
"name": "workspace",
|
||||||
|
"target": "/workspace",
|
||||||
|
"source_type": "repo",
|
||||||
|
"writable": True,
|
||||||
|
"owner": "user",
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "pi_state",
|
||||||
|
"target": "/tmp/.pi/agents",
|
||||||
|
"source_type": "instance",
|
||||||
|
"writable": True,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
"name": "pi_config",
|
||||||
|
"target": "/home/user/.pi",
|
||||||
|
"source_type": "git_mount",
|
||||||
|
"git_mount_ref": "dotfiles",
|
||||||
|
"writable": True,
|
||||||
|
"owner": "user",
|
||||||
|
},
|
||||||
|
],
|
||||||
|
"runtime": {
|
||||||
|
"command": ["/bin/bash"],
|
||||||
|
"stdin_open": True,
|
||||||
|
"tty": True,
|
||||||
|
"working_dir": "/workspace",
|
||||||
|
},
|
||||||
|
}
|
||||||
|
),
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
# ── Update existing pi-agent tool_type ───────────────────────────
|
||||||
|
conn.execute(
|
||||||
|
sa.text(
|
||||||
|
"""
|
||||||
|
UPDATE tool_types
|
||||||
|
SET manifest_id = :manifest_id,
|
||||||
|
definition_type = 'manifest',
|
||||||
|
dockerfile_template = NULL,
|
||||||
|
compose_template = NULL
|
||||||
|
WHERE name = 'pi-agent'
|
||||||
|
"""
|
||||||
|
),
|
||||||
|
{"manifest_id": PI_AGENT_MANIFEST_ID},
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def downgrade() -> None:
|
||||||
|
conn = op.get_bind()
|
||||||
|
|
||||||
|
# Restore pi-agent templates if manifest_id column exists
|
||||||
|
result = conn.execute(
|
||||||
|
sa.text("""
|
||||||
|
SELECT column_name FROM information_schema.columns
|
||||||
|
WHERE table_name = 'tool_types' AND column_name = 'manifest_id'
|
||||||
|
""")
|
||||||
|
)
|
||||||
|
has_manifest_id = result.fetchone() is not None
|
||||||
|
|
||||||
|
if has_manifest_id:
|
||||||
|
conn.execute(
|
||||||
|
sa.text(
|
||||||
|
"""
|
||||||
|
UPDATE tool_types
|
||||||
|
SET manifest_id = NULL,
|
||||||
|
definition_type = 'dockerfile',
|
||||||
|
dockerfile_template = :dockerfile,
|
||||||
|
compose_template = :compose
|
||||||
|
WHERE name = 'pi-agent'
|
||||||
|
"""
|
||||||
|
),
|
||||||
|
{
|
||||||
|
"dockerfile": """# Pi Coding Agent - Terminal-based coding harness
|
||||||
|
FROM ubuntu:24.04
|
||||||
|
|
||||||
|
ENV DEBIAN_FRONTEND=noninteractive
|
||||||
|
|
||||||
|
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/*
|
||||||
|
|
||||||
|
RUN curl -fsSL https://deb.nodesource.com/setup_20.x | bash - \\
|
||||||
|
&& apt-get install -y nodejs \\
|
||||||
|
&& rm -rf /var/lib/apt/lists/*
|
||||||
|
|
||||||
|
RUN npm install -g --ignore-scripts @earendil-works/pi-coding-agent
|
||||||
|
|
||||||
|
RUN useradd -m -s /bin/bash user
|
||||||
|
WORKDIR /home/user
|
||||||
|
|
||||||
|
RUN git config --global init.defaultBranch main \\
|
||||||
|
&& git config --global user.email "dev@headquarter.local" \\
|
||||||
|
&& git config --global user.name "Developer"
|
||||||
|
|
||||||
|
RUN echo 'set -g mouse on\\nset -g default-terminal "screen-256color"' > /home/user/.tmux.conf
|
||||||
|
|
||||||
|
RUN mkdir -p /home/user/.config/ranger \\
|
||||||
|
&& echo 'set preview_files true\\nset use_preview_script true' > /home/user/.config/ranger/rc.conf
|
||||||
|
|
||||||
|
RUN mkdir -p /home/user/.pi/agent
|
||||||
|
|
||||||
|
USER user
|
||||||
|
|
||||||
|
CMD ["/bin/bash"]
|
||||||
|
""",
|
||||||
|
"compose": """services:
|
||||||
|
app:
|
||||||
|
build: .
|
||||||
|
stdin_open: true
|
||||||
|
tty: true
|
||||||
|
volumes:
|
||||||
|
- ${REPO_PATH}:/workspace
|
||||||
|
working_dir: /workspace
|
||||||
|
command: /bin/bash""",
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
# Drop columns conditionally
|
||||||
|
result = conn.execute(
|
||||||
|
sa.text("""
|
||||||
|
SELECT column_name FROM information_schema.columns
|
||||||
|
WHERE table_name = 'tool_instances' AND column_name = 'image_tag'
|
||||||
|
""")
|
||||||
|
)
|
||||||
|
if result.fetchone():
|
||||||
|
op.drop_column("tool_instances", "image_tag")
|
||||||
|
|
||||||
|
result = conn.execute(
|
||||||
|
sa.text("""
|
||||||
|
SELECT column_name FROM information_schema.columns
|
||||||
|
WHERE table_name = 'tool_instances' AND column_name = 'manifest_compiled_at'
|
||||||
|
""")
|
||||||
|
)
|
||||||
|
if result.fetchone():
|
||||||
|
op.drop_column("tool_instances", "manifest_compiled_at")
|
||||||
|
|
||||||
|
if has_manifest_id:
|
||||||
|
op.drop_constraint(
|
||||||
|
"fk_tool_types_manifest_id", "tool_types", type_="foreignkey"
|
||||||
|
)
|
||||||
|
op.drop_column("tool_types", "manifest_id")
|
||||||
|
|
||||||
|
op.drop_table("tool_definition_manifests")
|
||||||
@@ -0,0 +1,89 @@
|
|||||||
|
"""drop tool_configs and config_folders tables
|
||||||
|
|
||||||
|
Revision ID: 2026_05_28_drop_tool_configs_and_config_folders
|
||||||
|
Revises: 2026_05_28_add_tool_definition_manifests
|
||||||
|
Create Date: 2026-05-28
|
||||||
|
|
||||||
|
"""
|
||||||
|
|
||||||
|
from typing import Sequence, Union
|
||||||
|
|
||||||
|
from alembic import op
|
||||||
|
import sqlalchemy as sa
|
||||||
|
|
||||||
|
# revision identifiers, used by Alembic.
|
||||||
|
revision: str = "2026_05_28_drop_tool_configs_and_config_folders"
|
||||||
|
down_revision: Union[str, None] = "2026_05_28_add_terminal_sessions"
|
||||||
|
branch_labels: Union[str, Sequence[str], None] = None
|
||||||
|
depends_on: Union[str, Sequence[str], None] = None
|
||||||
|
|
||||||
|
|
||||||
|
def upgrade() -> None:
|
||||||
|
conn = op.get_bind()
|
||||||
|
|
||||||
|
# Drop tool_configs table if it exists
|
||||||
|
result = conn.execute(
|
||||||
|
sa.text("""
|
||||||
|
SELECT table_name FROM information_schema.tables
|
||||||
|
WHERE table_name = 'tool_configs'
|
||||||
|
""")
|
||||||
|
)
|
||||||
|
if result.fetchone():
|
||||||
|
op.drop_table("tool_configs")
|
||||||
|
|
||||||
|
# Drop config_folders table if it exists
|
||||||
|
result = conn.execute(
|
||||||
|
sa.text("""
|
||||||
|
SELECT table_name FROM information_schema.tables
|
||||||
|
WHERE table_name = 'config_folders'
|
||||||
|
""")
|
||||||
|
)
|
||||||
|
if result.fetchone():
|
||||||
|
op.drop_table("config_folders")
|
||||||
|
|
||||||
|
|
||||||
|
def downgrade() -> None:
|
||||||
|
# Recreate config_folders table
|
||||||
|
op.create_table(
|
||||||
|
"config_folders",
|
||||||
|
sa.Column("id", sa.UUID(), nullable=False),
|
||||||
|
sa.Column("user_id", sa.UUID(), 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", sa.JSON(), default=dict, nullable=False),
|
||||||
|
sa.Column("project_overrides", sa.JSON(), default=dict, nullable=True),
|
||||||
|
sa.Column("is_active", sa.Boolean(), default=True, nullable=False),
|
||||||
|
sa.Column(
|
||||||
|
"created_at", sa.TIMESTAMP(timezone=True), server_default=sa.func.now()
|
||||||
|
),
|
||||||
|
sa.Column(
|
||||||
|
"updated_at", sa.TIMESTAMP(timezone=True), server_default=sa.func.now()
|
||||||
|
),
|
||||||
|
sa.PrimaryKeyConstraint("id"),
|
||||||
|
)
|
||||||
|
|
||||||
|
# Recreate tool_configs table
|
||||||
|
op.create_table(
|
||||||
|
"tool_configs",
|
||||||
|
sa.Column("id", sa.UUID(), nullable=False),
|
||||||
|
sa.Column("user_id", sa.UUID(), nullable=False),
|
||||||
|
sa.Column("tool_type_id", sa.UUID(), nullable=False),
|
||||||
|
sa.Column("project_id", sa.UUID(), nullable=True),
|
||||||
|
sa.Column("key", sa.String(255), nullable=False),
|
||||||
|
sa.Column("value", sa.Text(), nullable=False),
|
||||||
|
sa.Column("config_type", sa.String(20), default="env", nullable=False),
|
||||||
|
sa.Column("file_path", sa.String(1024), nullable=True),
|
||||||
|
sa.Column("port_override", sa.Integer(), nullable=True),
|
||||||
|
sa.Column("start_command", sa.Text(), nullable=True),
|
||||||
|
sa.Column("working_directory", sa.Text(), nullable=True),
|
||||||
|
sa.Column("environment_variables", sa.JSON(), default=dict, nullable=True),
|
||||||
|
sa.Column("volumes", sa.JSON(), default=list, nullable=True),
|
||||||
|
sa.Column(
|
||||||
|
"created_at", sa.TIMESTAMP(timezone=True), server_default=sa.func.now()
|
||||||
|
),
|
||||||
|
sa.Column(
|
||||||
|
"updated_at", sa.TIMESTAMP(timezone=True), server_default=sa.func.now()
|
||||||
|
),
|
||||||
|
sa.PrimaryKeyConstraint("id"),
|
||||||
|
)
|
||||||
@@ -0,0 +1,69 @@
|
|||||||
|
"""add notifications table
|
||||||
|
|
||||||
|
Revision ID: 2026_05_29_add_notifications_table
|
||||||
|
Revises: 2026_05_28_add_monitoring_tables
|
||||||
|
Create Date: 2026-05-29
|
||||||
|
|
||||||
|
"""
|
||||||
|
|
||||||
|
from collections.abc import Sequence
|
||||||
|
|
||||||
|
import sqlalchemy as sa
|
||||||
|
from alembic import op
|
||||||
|
|
||||||
|
# revision identifiers, used by Alembic.
|
||||||
|
revision: str = "2026_05_29_add_notifications_table"
|
||||||
|
down_revision: str | None = "2026_05_28_add_monitoring_tables"
|
||||||
|
branch_labels: str | Sequence[str] | None = None
|
||||||
|
depends_on: str | Sequence[str] | None = None
|
||||||
|
|
||||||
|
|
||||||
|
def upgrade() -> None:
|
||||||
|
op.create_table(
|
||||||
|
"notifications",
|
||||||
|
sa.Column("id", sa.Uuid(), nullable=False),
|
||||||
|
sa.Column("user_id", sa.Uuid(), nullable=False),
|
||||||
|
sa.Column("category", sa.String(length=32), nullable=False),
|
||||||
|
sa.Column("severity", sa.String(length=16), nullable=False),
|
||||||
|
sa.Column("title", sa.String(length=255), nullable=False),
|
||||||
|
sa.Column("message", sa.Text(), nullable=True),
|
||||||
|
sa.Column("source_type", sa.String(length=64), nullable=True),
|
||||||
|
sa.Column("source_id", sa.Uuid(), nullable=True),
|
||||||
|
sa.Column(
|
||||||
|
"metadata",
|
||||||
|
sa.JSON(),
|
||||||
|
nullable=False,
|
||||||
|
server_default="{}",
|
||||||
|
),
|
||||||
|
sa.Column("read_at", sa.DateTime(timezone=True), nullable=True),
|
||||||
|
sa.Column("dismissed_at", sa.DateTime(timezone=True), nullable=True),
|
||||||
|
sa.Column(
|
||||||
|
"created_at",
|
||||||
|
sa.DateTime(timezone=True),
|
||||||
|
server_default=sa.func.now(),
|
||||||
|
nullable=False,
|
||||||
|
),
|
||||||
|
sa.ForeignKeyConstraint(
|
||||||
|
["user_id"],
|
||||||
|
["users.id"],
|
||||||
|
ondelete="CASCADE",
|
||||||
|
),
|
||||||
|
sa.PrimaryKeyConstraint("id"),
|
||||||
|
)
|
||||||
|
op.create_index(
|
||||||
|
"idx_notifications_user_created_at",
|
||||||
|
"notifications",
|
||||||
|
["user_id", sa.text("created_at DESC")],
|
||||||
|
)
|
||||||
|
op.create_index(
|
||||||
|
"idx_notifications_user_unread",
|
||||||
|
"notifications",
|
||||||
|
["user_id", "read_at"],
|
||||||
|
postgresql_where=sa.text("read_at IS NULL"),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def downgrade() -> None:
|
||||||
|
op.drop_index("idx_notifications_user_unread", table_name="notifications")
|
||||||
|
op.drop_index("idx_notifications_user_created_at", table_name="notifications")
|
||||||
|
op.drop_table("notifications")
|
||||||
@@ -0,0 +1,27 @@
|
|||||||
|
"""add_ssh_key_ids_to_tool_instances
|
||||||
|
|
||||||
|
Revision ID: 2026_05_29_add_ssh_key_ids_to_tool_instances
|
||||||
|
Revises: 2026_05_29_drop_ssh_key_id_from_config_profiles
|
||||||
|
Create Date: 2026-05-29 12:46:00.000000
|
||||||
|
"""
|
||||||
|
|
||||||
|
from alembic import op
|
||||||
|
import sqlalchemy as sa
|
||||||
|
|
||||||
|
|
||||||
|
# revision identifiers, used by Alembic.
|
||||||
|
revision = "2026_05_29_add_ssh_key_ids_to_tool_instances"
|
||||||
|
down_revision = "2026_05_29_drop_ssh_key_id_from_config_profiles"
|
||||||
|
branch_labels = None
|
||||||
|
depends_on = None
|
||||||
|
|
||||||
|
|
||||||
|
def upgrade() -> None:
|
||||||
|
op.add_column(
|
||||||
|
"tool_instances",
|
||||||
|
sa.Column("ssh_key_ids", sa.JSON(), nullable=True),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def downgrade() -> None:
|
||||||
|
op.drop_column("tool_instances", "ssh_key_ids")
|
||||||
@@ -0,0 +1,32 @@
|
|||||||
|
"""drop_ssh_key_id_from_config_profiles
|
||||||
|
|
||||||
|
Revision ID: 2026_05_29_drop_ssh_key_id_from_config_profiles
|
||||||
|
Revises: 069d3da4dc9b
|
||||||
|
Create Date: 2026-05-29 12:45:00.000000
|
||||||
|
"""
|
||||||
|
|
||||||
|
from alembic import op
|
||||||
|
import sqlalchemy as sa
|
||||||
|
|
||||||
|
|
||||||
|
# revision identifiers, used by Alembic.
|
||||||
|
revision = "2026_05_29_drop_ssh_key_id_from_config_profiles"
|
||||||
|
down_revision = "069d3da4dc9b"
|
||||||
|
branch_labels = None
|
||||||
|
depends_on = None
|
||||||
|
|
||||||
|
|
||||||
|
def upgrade() -> None:
|
||||||
|
op.drop_column("config_profiles", "ssh_key_id")
|
||||||
|
|
||||||
|
|
||||||
|
def downgrade() -> None:
|
||||||
|
op.add_column(
|
||||||
|
"config_profiles",
|
||||||
|
sa.Column(
|
||||||
|
"ssh_key_id",
|
||||||
|
sa.Uuid(),
|
||||||
|
sa.ForeignKey("ssh_keys.id", ondelete="SET NULL"),
|
||||||
|
nullable=True,
|
||||||
|
),
|
||||||
|
)
|
||||||
@@ -0,0 +1,54 @@
|
|||||||
|
"""fix code-server bind-addr to host in DB template
|
||||||
|
|
||||||
|
Revision ID: 2026_05_29_fix_code_server_bind_addr
|
||||||
|
Revises: 2026_05_29_fix_web_tool_bind_address
|
||||||
|
Create Date: 2026-05-29 15:00:00.000000
|
||||||
|
|
||||||
|
"""
|
||||||
|
|
||||||
|
from typing import Sequence
|
||||||
|
|
||||||
|
from alembic import op
|
||||||
|
import sqlalchemy as sa
|
||||||
|
|
||||||
|
# revision identifiers, used by Alembic.
|
||||||
|
revision: str = "2026_05_29_fix_code_server_bind_addr"
|
||||||
|
down_revision: str | None = "2026_05_29_fix_web_tool_bind_address"
|
||||||
|
branch_labels: Sequence[str] | None = None
|
||||||
|
depends_on: Sequence[str] | None = None
|
||||||
|
|
||||||
|
|
||||||
|
def upgrade() -> None:
|
||||||
|
conn = op.get_bind()
|
||||||
|
|
||||||
|
# Find code-server tool types with broken --bind-addr in compose template
|
||||||
|
result = conn.execute(
|
||||||
|
sa.text("""
|
||||||
|
SELECT id, compose_template
|
||||||
|
FROM tool_types
|
||||||
|
WHERE name = 'code-server'
|
||||||
|
AND compose_template LIKE '%--bind-addr%'
|
||||||
|
""")
|
||||||
|
).fetchall()
|
||||||
|
|
||||||
|
for tool_id, compose_template in result:
|
||||||
|
updated = compose_template.replace(
|
||||||
|
"--bind-addr 0.0.0.0:8443", "--host 0.0.0.0"
|
||||||
|
).replace("--bind-addr", "--host 0.0.0.0")
|
||||||
|
|
||||||
|
conn.execute(
|
||||||
|
sa.text("""
|
||||||
|
UPDATE tool_types
|
||||||
|
SET compose_template = :compose_template
|
||||||
|
WHERE id = :id
|
||||||
|
"""),
|
||||||
|
{"compose_template": updated, "id": tool_id},
|
||||||
|
)
|
||||||
|
|
||||||
|
print(
|
||||||
|
f"Fixed code-server template ({tool_id}): replaced --bind-addr with --host"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def downgrade() -> None:
|
||||||
|
pass
|
||||||
@@ -0,0 +1,148 @@
|
|||||||
|
"""Fix code-server bind address to include port
|
||||||
|
|
||||||
|
Revision ID: 2026_05_29_fix_code_server_bind_addr_port
|
||||||
|
Revises: 2026_05_29_remove_lsio_command_override
|
||||||
|
Create Date: 2026-05-29 18:00:00.000000
|
||||||
|
|
||||||
|
"""
|
||||||
|
from typing import Sequence, Union
|
||||||
|
|
||||||
|
from alembic import op
|
||||||
|
import sqlalchemy as sa
|
||||||
|
import yaml
|
||||||
|
|
||||||
|
# revision identifiers, used by Alembic.
|
||||||
|
revision: str = "2026_05_29_fix_code_server_bind_addr_port"
|
||||||
|
down_revision: Union[str, None] = "2026_05_29_remove_lsio_command_override"
|
||||||
|
branch_labels: Union[str, Sequence[str], None] = None
|
||||||
|
depends_on: Union[str, Sequence[str], None] = None
|
||||||
|
|
||||||
|
|
||||||
|
def _fix_tool_type_templates(conn) -> None:
|
||||||
|
"""Fix code-server tool type templates with broken --host override."""
|
||||||
|
result = conn.execute(
|
||||||
|
sa.text("""
|
||||||
|
SELECT id, compose_template, default_port
|
||||||
|
FROM tool_types
|
||||||
|
WHERE name = 'code-server'
|
||||||
|
AND compose_template LIKE '%--host%'
|
||||||
|
""")
|
||||||
|
).fetchall()
|
||||||
|
|
||||||
|
for tool_id, compose_template, default_port in result:
|
||||||
|
port = default_port or 8443
|
||||||
|
expected = f"--bind-addr 0.0.0.0:{port}"
|
||||||
|
|
||||||
|
# Replace any line containing --host with the correct bind-addr
|
||||||
|
lines = compose_template.split("\n")
|
||||||
|
new_lines = []
|
||||||
|
modified = False
|
||||||
|
for line in lines:
|
||||||
|
if "command:" in line and "--host" in line:
|
||||||
|
indent = line[: len(line) - len(line.lstrip())]
|
||||||
|
new_lines.append(f"{indent}command: {expected}")
|
||||||
|
modified = True
|
||||||
|
else:
|
||||||
|
new_lines.append(line)
|
||||||
|
|
||||||
|
if not modified:
|
||||||
|
continue
|
||||||
|
|
||||||
|
updated = "\n".join(new_lines)
|
||||||
|
conn.execute(
|
||||||
|
sa.text("""
|
||||||
|
UPDATE tool_types
|
||||||
|
SET compose_template = :compose_template
|
||||||
|
WHERE id = :id
|
||||||
|
"""),
|
||||||
|
{"compose_template": updated, "id": tool_id},
|
||||||
|
)
|
||||||
|
print(f"Fixed code-server template ({tool_id}): replaced --host with {expected}")
|
||||||
|
|
||||||
|
|
||||||
|
def _fix_instance_compose_files(conn) -> None:
|
||||||
|
"""Fix existing instance compose files on disk with broken --host override."""
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
# Use information_schema to check if compose_path column exists
|
||||||
|
col_result = conn.execute(
|
||||||
|
sa.text("""
|
||||||
|
SELECT column_name
|
||||||
|
FROM information_schema.columns
|
||||||
|
WHERE table_name = 'tool_instances'
|
||||||
|
AND column_name = 'compose_path'
|
||||||
|
""")
|
||||||
|
).fetchone()
|
||||||
|
|
||||||
|
if not col_result:
|
||||||
|
print("compose_path column not found, skipping instance file fixes")
|
||||||
|
return
|
||||||
|
|
||||||
|
result = conn.execute(
|
||||||
|
sa.text("""
|
||||||
|
SELECT id, compose_path, tool_type_id
|
||||||
|
FROM tool_instances
|
||||||
|
WHERE compose_path IS NOT NULL
|
||||||
|
""")
|
||||||
|
).fetchall()
|
||||||
|
|
||||||
|
for instance_id, compose_path, tool_type_id in result:
|
||||||
|
path = Path(compose_path)
|
||||||
|
if not path.exists():
|
||||||
|
continue
|
||||||
|
|
||||||
|
try:
|
||||||
|
content = path.read_text()
|
||||||
|
except Exception:
|
||||||
|
continue
|
||||||
|
|
||||||
|
if "--host" not in content:
|
||||||
|
continue
|
||||||
|
|
||||||
|
# Get default_port from tool_type
|
||||||
|
port_result = conn.execute(
|
||||||
|
sa.text("""
|
||||||
|
SELECT default_port FROM tool_types WHERE id = :id
|
||||||
|
"""),
|
||||||
|
{"id": tool_type_id},
|
||||||
|
).fetchone()
|
||||||
|
port = port_result[0] if port_result and port_result[0] else 8443
|
||||||
|
expected = f"--bind-addr 0.0.0.0:{port}"
|
||||||
|
|
||||||
|
try:
|
||||||
|
data = yaml.safe_load(content)
|
||||||
|
except Exception:
|
||||||
|
continue
|
||||||
|
|
||||||
|
if not data or "services" not in data:
|
||||||
|
continue
|
||||||
|
|
||||||
|
modified = False
|
||||||
|
for svc in data["services"].values():
|
||||||
|
if "command" in svc:
|
||||||
|
cmd = svc["command"]
|
||||||
|
if "--host" in cmd:
|
||||||
|
svc["command"] = expected
|
||||||
|
modified = True
|
||||||
|
|
||||||
|
if not modified:
|
||||||
|
continue
|
||||||
|
|
||||||
|
try:
|
||||||
|
path.write_text(yaml.dump(data, default_flow_style=False))
|
||||||
|
print(
|
||||||
|
f"Fixed code-server instance compose ({instance_id}): "
|
||||||
|
f"replaced --host with {expected}"
|
||||||
|
)
|
||||||
|
except Exception as exc:
|
||||||
|
print(f"Failed to fix instance {instance_id}: {exc}")
|
||||||
|
|
||||||
|
|
||||||
|
def upgrade() -> None:
|
||||||
|
conn = op.get_bind()
|
||||||
|
_fix_tool_type_templates(conn)
|
||||||
|
_fix_instance_compose_files(conn)
|
||||||
|
|
||||||
|
|
||||||
|
def downgrade() -> None:
|
||||||
|
pass
|
||||||
@@ -0,0 +1,140 @@
|
|||||||
|
"""fix web tool bind address to 0.0.0.0
|
||||||
|
|
||||||
|
Revision ID: 2026_05_29_fix_web_tool_bind_address
|
||||||
|
Revises: 2026_05_29_remove_ssh_keys_mount_from_manifest
|
||||||
|
Create Date: 2026-05-29 14: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_29_fix_web_tool_bind_address"
|
||||||
|
down_revision: Union[str, None] = "2026_05_29_remove_ssh_keys_mount_from_manifest"
|
||||||
|
branch_labels: Union[str, Sequence[str], None] = None
|
||||||
|
depends_on: Union[str, Sequence[str], None] = None
|
||||||
|
|
||||||
|
|
||||||
|
def _fix_code_server_compose(conn) -> None:
|
||||||
|
"""Update code-server compose template to bind to 0.0.0.0."""
|
||||||
|
result = conn.execute(
|
||||||
|
sa.text("""
|
||||||
|
SELECT id, compose_template, definition_type
|
||||||
|
FROM tool_types
|
||||||
|
WHERE name = 'code-server'
|
||||||
|
""")
|
||||||
|
).fetchone()
|
||||||
|
|
||||||
|
if result is None:
|
||||||
|
return
|
||||||
|
|
||||||
|
tool_id, compose_template, definition_type = result
|
||||||
|
|
||||||
|
if definition_type != "compose" or not compose_template:
|
||||||
|
return
|
||||||
|
|
||||||
|
# Fix or add command to bind to 0.0.0.0
|
||||||
|
lines = compose_template.split("\n")
|
||||||
|
new_lines = []
|
||||||
|
image_line_idx = -1
|
||||||
|
command_fixed = False
|
||||||
|
for i, line in enumerate(lines):
|
||||||
|
# Replace broken --bind-addr with correct --host
|
||||||
|
if "command:" in line and "--bind-addr" in line:
|
||||||
|
indent = line[: len(line) - len(line.lstrip())]
|
||||||
|
new_lines.append(f"{indent}command: --host 0.0.0.0")
|
||||||
|
command_fixed = True
|
||||||
|
continue
|
||||||
|
new_lines.append(line)
|
||||||
|
if "image:" in line and image_line_idx == -1:
|
||||||
|
image_line_idx = i
|
||||||
|
|
||||||
|
# If no command line exists, insert one after image
|
||||||
|
if not command_fixed and image_line_idx != -1:
|
||||||
|
image_line = lines[image_line_idx]
|
||||||
|
indent = image_line[: len(image_line) - len(image_line.lstrip())]
|
||||||
|
# Insert after the image line in new_lines
|
||||||
|
insert_idx = new_lines.index(image_line) + 1
|
||||||
|
new_lines.insert(insert_idx, f"{indent}command: --host 0.0.0.0")
|
||||||
|
command_fixed = True
|
||||||
|
|
||||||
|
if not command_fixed:
|
||||||
|
return
|
||||||
|
|
||||||
|
updated_compose = "\n".join(new_lines)
|
||||||
|
|
||||||
|
conn.execute(
|
||||||
|
sa.text("""
|
||||||
|
UPDATE tool_types
|
||||||
|
SET compose_template = :compose_template
|
||||||
|
WHERE id = :id
|
||||||
|
"""),
|
||||||
|
{"compose_template": updated_compose, "id": tool_id},
|
||||||
|
)
|
||||||
|
|
||||||
|
print(f"Updated code-server tool type ({tool_id}) to bind to 0.0.0.0")
|
||||||
|
|
||||||
|
|
||||||
|
def _fix_jupyter_compose(conn) -> None:
|
||||||
|
"""Update jupyter-notebook compose template to bind to 0.0.0.0."""
|
||||||
|
result = conn.execute(
|
||||||
|
sa.text("""
|
||||||
|
SELECT id, compose_template, definition_type
|
||||||
|
FROM tool_types
|
||||||
|
WHERE name = 'jupyter-notebook'
|
||||||
|
""")
|
||||||
|
).fetchone()
|
||||||
|
|
||||||
|
if result is None:
|
||||||
|
return
|
||||||
|
|
||||||
|
tool_id, compose_template, definition_type = result
|
||||||
|
|
||||||
|
if definition_type != "compose" or not compose_template:
|
||||||
|
return
|
||||||
|
|
||||||
|
if "command:" in compose_template:
|
||||||
|
return
|
||||||
|
|
||||||
|
lines = compose_template.split("\n")
|
||||||
|
new_lines = []
|
||||||
|
image_line_idx = -1
|
||||||
|
for i, line in enumerate(lines):
|
||||||
|
new_lines.append(line)
|
||||||
|
if "image:" in line and image_line_idx == -1:
|
||||||
|
image_line_idx = i
|
||||||
|
indent = line[: len(line) - len(line.lstrip())]
|
||||||
|
# Jupyter needs --ip=0.0.0.0 to bind to all interfaces
|
||||||
|
new_lines.append(
|
||||||
|
f"{indent}command: start-notebook.sh --ip=0.0.0.0 --port=8888 --no-browser"
|
||||||
|
)
|
||||||
|
|
||||||
|
if image_line_idx == -1:
|
||||||
|
return
|
||||||
|
|
||||||
|
updated_compose = "\n".join(new_lines)
|
||||||
|
|
||||||
|
conn.execute(
|
||||||
|
sa.text("""
|
||||||
|
UPDATE tool_types
|
||||||
|
SET compose_template = :compose_template
|
||||||
|
WHERE id = :id
|
||||||
|
"""),
|
||||||
|
{"compose_template": updated_compose, "id": tool_id},
|
||||||
|
)
|
||||||
|
|
||||||
|
print(f"Updated jupyter-notebook tool type ({tool_id}) to bind to 0.0.0.0:8888")
|
||||||
|
|
||||||
|
|
||||||
|
def upgrade() -> None:
|
||||||
|
conn = op.get_bind()
|
||||||
|
_fix_code_server_compose(conn)
|
||||||
|
_fix_jupyter_compose(conn)
|
||||||
|
|
||||||
|
|
||||||
|
def downgrade() -> None:
|
||||||
|
# Cannot safely downgrade without knowing the original compose_template
|
||||||
|
pass
|
||||||
@@ -0,0 +1,121 @@
|
|||||||
|
"""Remove broken command override from LSIO code-server templates
|
||||||
|
|
||||||
|
Revision ID: 2026_05_29_remove_lsio_command_override
|
||||||
|
Revises: 2026_05_29_fix_code_server_bind_addr
|
||||||
|
Create Date: 2026-05-29 15:05:00.000000
|
||||||
|
|
||||||
|
"""
|
||||||
|
|
||||||
|
from collections.abc import Sequence
|
||||||
|
|
||||||
|
from alembic import op
|
||||||
|
import sqlalchemy as sa
|
||||||
|
|
||||||
|
# revision identifiers, used by Alembic.
|
||||||
|
revision: str = "2026_05_29_remove_lsio_command_override"
|
||||||
|
down_revision: str | None = "2026_05_29_fix_code_server_bind_addr"
|
||||||
|
branch_labels: Sequence[str] | None = None
|
||||||
|
depends_on: Sequence[str] | None = None
|
||||||
|
|
||||||
|
|
||||||
|
def upgrade() -> None:
|
||||||
|
conn = op.get_bind()
|
||||||
|
|
||||||
|
# Fix tool_types templates in DB
|
||||||
|
result = conn.execute(
|
||||||
|
sa.text("""
|
||||||
|
SELECT id, compose_template
|
||||||
|
FROM tool_types
|
||||||
|
WHERE name = 'code-server'
|
||||||
|
""")
|
||||||
|
).fetchall()
|
||||||
|
|
||||||
|
import yaml
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
for tool_id, compose_template in result:
|
||||||
|
try:
|
||||||
|
data = yaml.safe_load(compose_template)
|
||||||
|
except Exception:
|
||||||
|
continue
|
||||||
|
|
||||||
|
if not data or "services" not in data:
|
||||||
|
continue
|
||||||
|
|
||||||
|
modified = False
|
||||||
|
for svc in data["services"].values():
|
||||||
|
image = svc.get("image", "")
|
||||||
|
if not image or "linuxserver" not in image:
|
||||||
|
continue
|
||||||
|
if "command" in svc:
|
||||||
|
cmd = svc["command"]
|
||||||
|
if "--bind-addr" in cmd or "--host" in cmd:
|
||||||
|
del svc["command"]
|
||||||
|
modified = True
|
||||||
|
|
||||||
|
if modified:
|
||||||
|
updated = yaml.dump(data, default_flow_style=False)
|
||||||
|
conn.execute(
|
||||||
|
sa.text("""
|
||||||
|
UPDATE tool_types
|
||||||
|
SET compose_template = :compose_template
|
||||||
|
WHERE id = :id
|
||||||
|
"""),
|
||||||
|
{"compose_template": updated, "id": tool_id},
|
||||||
|
)
|
||||||
|
print(f"Removed broken command override from LSIO template ({tool_id})")
|
||||||
|
|
||||||
|
# Fix existing instance compose files on disk
|
||||||
|
# Use information_schema to check if compose_path column exists
|
||||||
|
col_result = conn.execute(
|
||||||
|
sa.text("""
|
||||||
|
SELECT column_name
|
||||||
|
FROM information_schema.columns
|
||||||
|
WHERE table_name = 'tool_instances'
|
||||||
|
AND column_name = 'compose_path'
|
||||||
|
""")
|
||||||
|
).fetchone()
|
||||||
|
|
||||||
|
if col_result:
|
||||||
|
result = conn.execute(
|
||||||
|
sa.text("""
|
||||||
|
SELECT id, compose_path
|
||||||
|
FROM tool_instances
|
||||||
|
WHERE compose_path IS NOT NULL
|
||||||
|
""")
|
||||||
|
).fetchall()
|
||||||
|
|
||||||
|
for instance_id, compose_path in result:
|
||||||
|
path = Path(compose_path)
|
||||||
|
if not path.exists():
|
||||||
|
continue
|
||||||
|
try:
|
||||||
|
content = path.read_text()
|
||||||
|
data = yaml.safe_load(content)
|
||||||
|
except Exception:
|
||||||
|
continue
|
||||||
|
|
||||||
|
if not data or "services" not in data:
|
||||||
|
continue
|
||||||
|
|
||||||
|
modified = False
|
||||||
|
for svc in data["services"].values():
|
||||||
|
image = svc.get("image", "")
|
||||||
|
if not image or "linuxserver" not in image:
|
||||||
|
continue
|
||||||
|
if "command" in svc:
|
||||||
|
cmd = svc["command"]
|
||||||
|
if "--bind-addr" in cmd or "--host" in cmd:
|
||||||
|
del svc["command"]
|
||||||
|
modified = True
|
||||||
|
|
||||||
|
if modified:
|
||||||
|
path.write_text(yaml.dump(data, default_flow_style=False))
|
||||||
|
print(
|
||||||
|
f"Removed broken command override from instance compose "
|
||||||
|
f"({instance_id})"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def downgrade() -> None:
|
||||||
|
pass
|
||||||
@@ -0,0 +1,105 @@
|
|||||||
|
"""remove ssh_keys mount from pi-agent manifest
|
||||||
|
|
||||||
|
Revision ID: 2026_05_29_remove_ssh_keys_mount_from_manifest
|
||||||
|
Revises: 2026_05_29_add_ssh_key_ids_to_tool_instances
|
||||||
|
Create Date: 2026-05-29 14:00:00.000000
|
||||||
|
|
||||||
|
"""
|
||||||
|
|
||||||
|
import json
|
||||||
|
from typing import Sequence, Union
|
||||||
|
|
||||||
|
from alembic import op
|
||||||
|
import sqlalchemy as sa
|
||||||
|
|
||||||
|
# revision identifiers, used by Alembic.
|
||||||
|
revision: str = "2026_05_29_remove_ssh_keys_mount_from_manifest"
|
||||||
|
down_revision: Union[str, None] = "2026_05_29_add_ssh_key_ids_to_tool_instances"
|
||||||
|
branch_labels: Union[str, Sequence[str], None] = None
|
||||||
|
depends_on: Union[str, Sequence[str], None] = None
|
||||||
|
|
||||||
|
|
||||||
|
def upgrade() -> None:
|
||||||
|
"""Remove the ssh_keys mount from the pi-agent manifest."""
|
||||||
|
conn = op.get_bind()
|
||||||
|
|
||||||
|
# Get the pi-agent manifest
|
||||||
|
result = conn.execute(
|
||||||
|
sa.text(
|
||||||
|
"SELECT id, manifest FROM tool_definition_manifests WHERE name = 'pi-agent'"
|
||||||
|
)
|
||||||
|
)
|
||||||
|
row = result.fetchone()
|
||||||
|
if not row:
|
||||||
|
return
|
||||||
|
|
||||||
|
manifest_id, manifest_json = row
|
||||||
|
manifest = (
|
||||||
|
manifest_json if isinstance(manifest_json, dict) else json.loads(manifest_json)
|
||||||
|
)
|
||||||
|
|
||||||
|
mounts = manifest.get("mounts", [])
|
||||||
|
original_count = len(mounts)
|
||||||
|
|
||||||
|
# Remove any mount named "ssh_keys"
|
||||||
|
filtered_mounts = [m for m in mounts if m.get("name") != "ssh_keys"]
|
||||||
|
|
||||||
|
if len(filtered_mounts) < original_count:
|
||||||
|
manifest["mounts"] = filtered_mounts
|
||||||
|
conn.execute(
|
||||||
|
sa.text(
|
||||||
|
"UPDATE tool_definition_manifests SET manifest = :manifest WHERE id = :id"
|
||||||
|
),
|
||||||
|
{
|
||||||
|
"manifest": json.dumps(manifest),
|
||||||
|
"id": manifest_id,
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def downgrade() -> None:
|
||||||
|
"""Restore the ssh_keys mount to the pi-agent manifest."""
|
||||||
|
conn = op.get_bind()
|
||||||
|
|
||||||
|
result = conn.execute(
|
||||||
|
sa.text(
|
||||||
|
"SELECT id, manifest FROM tool_definition_manifests WHERE name = 'pi-agent'"
|
||||||
|
)
|
||||||
|
)
|
||||||
|
row = result.fetchone()
|
||||||
|
if not row:
|
||||||
|
return
|
||||||
|
|
||||||
|
manifest_id, manifest_json = row
|
||||||
|
manifest = (
|
||||||
|
manifest_json if isinstance(manifest_json, dict) else json.loads(manifest_json)
|
||||||
|
)
|
||||||
|
|
||||||
|
mounts = manifest.get("mounts", [])
|
||||||
|
|
||||||
|
# Check if ssh_keys mount already exists
|
||||||
|
if any(m.get("name") == "ssh_keys" for m in mounts):
|
||||||
|
return
|
||||||
|
|
||||||
|
# Add the ssh_keys mount back
|
||||||
|
mounts.append(
|
||||||
|
{
|
||||||
|
"name": "ssh_keys",
|
||||||
|
"target": "/home/user/.ssh",
|
||||||
|
"source_type": "ssh_key",
|
||||||
|
"mode": "0700",
|
||||||
|
"file_mode": "0600",
|
||||||
|
"readonly": True,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
manifest["mounts"] = mounts
|
||||||
|
|
||||||
|
conn.execute(
|
||||||
|
sa.text(
|
||||||
|
"UPDATE tool_definition_manifests SET manifest = :manifest WHERE id = :id"
|
||||||
|
),
|
||||||
|
{
|
||||||
|
"manifest": json.dumps(manifest),
|
||||||
|
"id": manifest_id,
|
||||||
|
},
|
||||||
|
)
|
||||||
@@ -5,8 +5,6 @@ Revises: 2026_05_23_remove_is_builtin, 2026_05_24_add_config_profiles
|
|||||||
Create Date: 2026-05-24 18:00:43.990361
|
Create Date: 2026-05-24 18:00:43.990361
|
||||||
"""
|
"""
|
||||||
|
|
||||||
from alembic import op
|
|
||||||
import sqlalchemy as sa
|
|
||||||
|
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -7,8 +7,6 @@ Create Date: 2026-05-24 10:43:14.000000
|
|||||||
"""
|
"""
|
||||||
from typing import Sequence, Union
|
from typing import Sequence, Union
|
||||||
|
|
||||||
from alembic import op
|
|
||||||
import sqlalchemy as sa
|
|
||||||
|
|
||||||
# revision identifiers, used by Alembic.
|
# revision identifiers, used by Alembic.
|
||||||
revision: str = "f3d2dc90ba3a"
|
revision: str = "f3d2dc90ba3a"
|
||||||
|
|||||||
@@ -1,4 +1,6 @@
|
|||||||
from src.api.auth import router as auth_router
|
from src.api.auth import router as auth_router
|
||||||
|
from src.api.events import router as events_router
|
||||||
|
from src.api.notifications import router as notifications_router
|
||||||
from src.api.users import router as users_router
|
from src.api.users import router as users_router
|
||||||
|
|
||||||
__all__ = ["auth_router", "users_router"]
|
__all__ = ["auth_router", "events_router", "notifications_router", "users_router"]
|
||||||
|
|||||||
+10
-10
@@ -48,7 +48,7 @@ async def login(next: str = "/") -> RedirectResponse:
|
|||||||
redirect_uri=redirect_uri,
|
redirect_uri=redirect_uri,
|
||||||
state=state,
|
state=state,
|
||||||
)
|
)
|
||||||
logger.info("Auth login initiated: redirect_uri=%s, next=%s", redirect_uri, next)
|
logger.debug("Auth login initiated: redirect_uri=%s, next=%s", redirect_uri, next)
|
||||||
response = RedirectResponse(location)
|
response = RedirectResponse(location)
|
||||||
response.set_cookie("auth_state", state, httponly=True, samesite="lax")
|
response.set_cookie("auth_state", state, httponly=True, samesite="lax")
|
||||||
response.set_cookie("auth_next", next, httponly=True, samesite="lax")
|
response.set_cookie("auth_next", next, httponly=True, samesite="lax")
|
||||||
@@ -63,7 +63,7 @@ async def callback(
|
|||||||
auth_next: str | None = Cookie(default="/"),
|
auth_next: str | None = Cookie(default="/"),
|
||||||
session: AsyncSession = Depends(get_db_session),
|
session: AsyncSession = Depends(get_db_session),
|
||||||
) -> RedirectResponse:
|
) -> RedirectResponse:
|
||||||
logger.info("Auth callback received: code=%s... state=%s", code[:10] if code else "None", state[:10] if state else "None")
|
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:
|
if auth_state is None or auth_state != state:
|
||||||
logger.warning("State mismatch: cookie=%s, param=%s", auth_state, state)
|
logger.warning("State mismatch: cookie=%s, param=%s", auth_state, state)
|
||||||
@@ -71,7 +71,7 @@ async def callback(
|
|||||||
|
|
||||||
settings = Settings()
|
settings = Settings()
|
||||||
redirect_uri = f"{settings.api_base_url}/auth/callback"
|
redirect_uri = f"{settings.api_base_url}/auth/callback"
|
||||||
logger.info("Exchanging code for tokens (redirect_uri=%s)", redirect_uri)
|
logger.debug("Exchanging code for tokens (redirect_uri=%s)", redirect_uri)
|
||||||
|
|
||||||
async with httpx.AsyncClient() as client:
|
async with httpx.AsyncClient() as client:
|
||||||
try:
|
try:
|
||||||
@@ -92,7 +92,7 @@ async def callback(
|
|||||||
access_token=token_payload["access_token"],
|
access_token=token_payload["access_token"],
|
||||||
client=client,
|
client=client,
|
||||||
)
|
)
|
||||||
logger.info("User info fetched successfully")
|
logger.debug("User info fetched successfully")
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
logger.error("User info fetch failed: %s", 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")
|
raise HTTPException(status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail="failed to fetch user info")
|
||||||
@@ -100,19 +100,19 @@ async def callback(
|
|||||||
authentik_id = str(user_info.get("sub", ""))
|
authentik_id = str(user_info.get("sub", ""))
|
||||||
email = str(user_info.get("email", f"{authentik_id}@authentik.local"))
|
email = str(user_info.get("email", f"{authentik_id}@authentik.local"))
|
||||||
name = str(user_info.get("name", email))
|
name = str(user_info.get("name", email))
|
||||||
logger.info("User info: authentik_id=%s, email=%s, name=%s", authentik_id, email, name)
|
logger.debug("User info: authentik_id=%s, email=%s, name=%s", authentik_id, email, name)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
user = await session.scalar(select(User).where(User.authentik_id == authentik_id))
|
user = await session.scalar(select(User).where(User.authentik_id == authentik_id))
|
||||||
if user is None:
|
if user is None:
|
||||||
logger.info("Creating new user: authentik_id=%s", authentik_id)
|
logger.debug("Creating new user: authentik_id=%s", authentik_id)
|
||||||
user = User(email=email, name=name, authentik_id=authentik_id, avatar_url=None)
|
user = User(email=email, name=name, authentik_id=authentik_id, avatar_url=None)
|
||||||
session.add(user)
|
session.add(user)
|
||||||
await session.commit()
|
await session.commit()
|
||||||
await session.refresh(user)
|
await session.refresh(user)
|
||||||
logger.info("New user created: id=%s", user.id)
|
logger.info("New user created: id=%s", user.id)
|
||||||
else:
|
else:
|
||||||
logger.info("Existing user found: id=%s, updating info", user.id)
|
logger.debug("Existing user found: id=%s, updating info", user.id)
|
||||||
user.email = email
|
user.email = email
|
||||||
user.name = name
|
user.name = name
|
||||||
await session.commit()
|
await session.commit()
|
||||||
@@ -165,20 +165,20 @@ async def me(
|
|||||||
session_cookie: str | None = Cookie(default=None, alias="session"),
|
session_cookie: str | None = Cookie(default=None, alias="session"),
|
||||||
session: AsyncSession = Depends(get_db_session),
|
session: AsyncSession = Depends(get_db_session),
|
||||||
) -> dict[str, Any]:
|
) -> dict[str, Any]:
|
||||||
logger.info("Auth /me called, cookie present: %s", bool(session_cookie))
|
logger.debug("Auth /me called, cookie present: %s", bool(session_cookie))
|
||||||
|
|
||||||
if not session_cookie:
|
if not session_cookie:
|
||||||
logger.warning("Auth /me: missing session cookie")
|
logger.warning("Auth /me: missing session cookie")
|
||||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="missing session")
|
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="missing session")
|
||||||
|
|
||||||
settings = Settings()
|
settings = Settings()
|
||||||
logger.info("Auth /me: cookie_domain=%s, cookie_secure=%s, cookie_samesite=%s",
|
logger.debug("Auth /me: cookie_domain=%s, cookie_secure=%s, cookie_samesite=%s",
|
||||||
settings.cookie_domain, settings.cookie_secure, settings.cookie_samesite)
|
settings.cookie_domain, settings.cookie_secure, settings.cookie_samesite)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
payload = decode_session_cookie(settings=settings, cookie_value=session_cookie)
|
payload = decode_session_cookie(settings=settings, cookie_value=session_cookie)
|
||||||
user_id = payload["user_id"]
|
user_id = payload["user_id"]
|
||||||
logger.info("Auth /me: decoded session for user_id=%s", user_id)
|
logger.debug("Auth /me: decoded session for user_id=%s", user_id)
|
||||||
except ValueError as exc:
|
except ValueError as exc:
|
||||||
logger.warning("Auth /me: invalid session: %s", exc)
|
logger.warning("Auth /me: invalid session: %s", exc)
|
||||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail=str(exc))
|
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail=str(exc))
|
||||||
|
|||||||
@@ -1,372 +0,0 @@
|
|||||||
"""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.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"])
|
|
||||||
|
|
||||||
MAX_FOLDER_SIZE_MB = 10
|
|
||||||
MAX_FOLDER_SIZE_BYTES = MAX_FOLDER_SIZE_MB * 1024 * 1024
|
|
||||||
|
|
||||||
|
|
||||||
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:
|
|
||||||
if not v.startswith("/"):
|
|
||||||
raise ValueError("Mount path must be absolute (start with /)")
|
|
||||||
return v
|
|
||||||
|
|
||||||
@field_validator("files")
|
|
||||||
@classmethod
|
|
||||||
def validate_files(cls, v: dict) -> dict:
|
|
||||||
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_FOLDER_SIZE_BYTES:
|
|
||||||
raise ValueError(f"Total folder size exceeds {MAX_FOLDER_SIZE_MB}MB limit")
|
|
||||||
|
|
||||||
return 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:
|
|
||||||
if v is None:
|
|
||||||
return v
|
|
||||||
if not v.startswith("/"):
|
|
||||||
raise ValueError("Mount path must be absolute (start with /)")
|
|
||||||
return v
|
|
||||||
|
|
||||||
@field_validator("files")
|
|
||||||
@classmethod
|
|
||||||
def validate_files(cls, v: dict | None) -> dict | None:
|
|
||||||
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_FOLDER_SIZE_BYTES:
|
|
||||||
raise ValueError(f"Total folder size exceeds {MAX_FOLDER_SIZE_MB}MB limit")
|
|
||||||
|
|
||||||
return 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:
|
|
||||||
if v is None:
|
|
||||||
return v
|
|
||||||
if not v.startswith("/"):
|
|
||||||
raise ValueError("Mount path must be absolute (start with /)")
|
|
||||||
return 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 {},
|
|
||||||
}
|
|
||||||
@@ -1,14 +1,18 @@
|
|||||||
"""Config profile API endpoints."""
|
"""Config profile API endpoints."""
|
||||||
|
|
||||||
import logging
|
import logging
|
||||||
|
import os
|
||||||
|
import subprocess
|
||||||
import uuid
|
import uuid
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
from fastapi import APIRouter, Depends, HTTPException, Query, status
|
from fastapi import APIRouter, Depends, HTTPException, Query, status
|
||||||
from pydantic import BaseModel, Field, field_validator
|
from pydantic import BaseModel, Field, field_validator, model_validator
|
||||||
from sqlalchemy import select
|
from sqlalchemy import select
|
||||||
from sqlalchemy.ext.asyncio import AsyncSession
|
from sqlalchemy.ext.asyncio import AsyncSession
|
||||||
from sqlalchemy.orm import selectinload
|
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.auth.dependencies import get_current_user_id, get_db_session
|
||||||
from src.models.config_profile import ConfigProfile, ConfigProfileInclude
|
from src.models.config_profile import ConfigProfile, ConfigProfileInclude
|
||||||
from src.models.project import Project
|
from src.models.project import Project
|
||||||
@@ -19,6 +23,7 @@ from src.services.config_profile_resolver import (
|
|||||||
resolve_profile,
|
resolve_profile,
|
||||||
resolved_profile_to_dict,
|
resolved_profile_to_dict,
|
||||||
)
|
)
|
||||||
|
from src.utils.git_url_parser import parse_git_url
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
@@ -55,10 +60,89 @@ def _calculate_profile_size(data: dict) -> int:
|
|||||||
return total
|
return total
|
||||||
|
|
||||||
|
|
||||||
|
class GitMountMapping(BaseModel):
|
||||||
|
source_path: str = Field(
|
||||||
|
description="Path within repository (supports glob patterns)"
|
||||||
|
)
|
||||||
|
target_path: str = Field(description="Absolute path inside container")
|
||||||
|
|
||||||
|
@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 GitMountItem(BaseModel):
|
||||||
|
remote_url: str = Field(description="Git remote URL (HTTPS or SSH)")
|
||||||
|
source_path: str | None = Field(
|
||||||
|
default=None, description="Path within repository (legacy single mapping)"
|
||||||
|
)
|
||||||
|
target_path: str | None = Field(
|
||||||
|
default=None,
|
||||||
|
description="Absolute path inside container (legacy single mapping)",
|
||||||
|
)
|
||||||
|
branch: str | None = Field(default=None, description="Optional branch or tag name")
|
||||||
|
mappings: list[GitMountMapping] | None = Field(
|
||||||
|
default=None, description="Multiple source/target mappings from the same repo"
|
||||||
|
)
|
||||||
|
|
||||||
|
@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 | None) -> str | None:
|
||||||
|
if v is None:
|
||||||
|
return v
|
||||||
|
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 | None) -> str | None:
|
||||||
|
if v is None:
|
||||||
|
return v
|
||||||
|
if ".." in v:
|
||||||
|
raise ValueError("target_path cannot contain path traversal (..)")
|
||||||
|
return v
|
||||||
|
|
||||||
|
@model_validator(mode="after")
|
||||||
|
def check_mappings_or_legacy(self):
|
||||||
|
has_legacy = self.source_path is not None and self.target_path is not None
|
||||||
|
has_mappings = self.mappings is not None and len(self.mappings) > 0
|
||||||
|
if not has_legacy and not has_mappings:
|
||||||
|
raise ValueError(
|
||||||
|
"Git mount must have either 'mappings' (non-empty array) or both 'source_path' and 'target_path'"
|
||||||
|
)
|
||||||
|
return self
|
||||||
|
|
||||||
|
|
||||||
class MountItem(BaseModel):
|
class MountItem(BaseModel):
|
||||||
target: str = Field(description="Absolute mount target path")
|
target: str = Field(description="Absolute mount target path")
|
||||||
mode: str = Field(default="rw", description="Mount mode: ro or rw")
|
mode: str = Field(default="rw", description="Mount mode: ro or rw")
|
||||||
files: dict = Field(default_factory=dict, description="Files as {relative_path: content}")
|
files: dict = Field(
|
||||||
|
default_factory=dict, description="Files as {relative_path: content}"
|
||||||
|
)
|
||||||
|
|
||||||
@field_validator("target")
|
@field_validator("target")
|
||||||
@classmethod
|
@classmethod
|
||||||
@@ -95,9 +179,18 @@ class ConfigProfileCreate(BaseModel):
|
|||||||
tool_type_id: str | None = Field(default=None, description="Optional tool type ID")
|
tool_type_id: str | None = Field(default=None, description="Optional tool type ID")
|
||||||
env_vars: dict = Field(default_factory=dict, description="Environment variables")
|
env_vars: dict = Field(default_factory=dict, description="Environment variables")
|
||||||
runtime_hints: dict = Field(default_factory=dict, description="Runtime hints")
|
runtime_hints: dict = Field(default_factory=dict, description="Runtime hints")
|
||||||
mounts: list[MountItem] = Field(default_factory=list, description="Mount definitions")
|
mounts: list[MountItem] = Field(
|
||||||
files: dict = Field(default_factory=dict, description="Files as {relative_path: content}")
|
default_factory=list, description="Mount definitions"
|
||||||
is_default: bool = Field(default=False, description="Whether this is the default profile for its scope")
|
)
|
||||||
|
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")
|
@field_validator("project_id", "tool_type_id")
|
||||||
@classmethod
|
@classmethod
|
||||||
@@ -120,9 +213,10 @@ class ConfigProfileCreate(BaseModel):
|
|||||||
@field_validator("env_vars")
|
@field_validator("env_vars")
|
||||||
@classmethod
|
@classmethod
|
||||||
def validate_env_vars(cls, v: dict) -> dict:
|
def validate_env_vars(cls, v: dict) -> dict:
|
||||||
if not isinstance(v, dict):
|
result = _validate_env_vars(v)
|
||||||
|
if result is None:
|
||||||
raise ValueError("env_vars must be a JSON object")
|
raise ValueError("env_vars must be a JSON object")
|
||||||
return v
|
return result
|
||||||
|
|
||||||
@field_validator("runtime_hints")
|
@field_validator("runtime_hints")
|
||||||
@classmethod
|
@classmethod
|
||||||
@@ -146,9 +240,18 @@ class ConfigProfileUpdate(BaseModel):
|
|||||||
tool_type_id: str | None = Field(default=None, description="Optional tool type ID")
|
tool_type_id: str | None = Field(default=None, description="Optional tool type ID")
|
||||||
env_vars: dict | None = Field(default=None, description="Environment variables")
|
env_vars: dict | None = Field(default=None, description="Environment variables")
|
||||||
runtime_hints: dict | None = Field(default=None, description="Runtime hints")
|
runtime_hints: dict | None = Field(default=None, description="Runtime hints")
|
||||||
mounts: list[MountItem] | None = Field(default=None, description="Mount definitions")
|
mounts: list[MountItem] | None = Field(
|
||||||
files: dict | None = Field(default=None, description="Files as {relative_path: content}")
|
default=None, description="Mount definitions"
|
||||||
is_default: bool | None = Field(default=None, description="Whether this is the default profile")
|
)
|
||||||
|
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")
|
@field_validator("project_id", "tool_type_id")
|
||||||
@classmethod
|
@classmethod
|
||||||
@@ -191,13 +294,16 @@ class ConfigProfileResponse(BaseModel):
|
|||||||
runtime_hints: dict
|
runtime_hints: dict
|
||||||
mounts: list
|
mounts: list
|
||||||
files: dict
|
files: dict
|
||||||
|
git_mounts: list
|
||||||
is_default: bool
|
is_default: bool
|
||||||
includes: list[dict]
|
includes: list[dict]
|
||||||
created_at: str
|
created_at: str
|
||||||
updated_at: str
|
updated_at: str
|
||||||
|
|
||||||
|
|
||||||
async def _get_profile_with_includes(session: AsyncSession, profile_id: uuid.UUID) -> ConfigProfile | None:
|
async def _get_profile_with_includes(
|
||||||
|
session: AsyncSession, profile_id: uuid.UUID
|
||||||
|
) -> ConfigProfile | None:
|
||||||
"""Fetch a profile with includes eagerly loaded."""
|
"""Fetch a profile with includes eagerly loaded."""
|
||||||
result = await session.execute(
|
result = await session.execute(
|
||||||
select(ConfigProfile)
|
select(ConfigProfile)
|
||||||
@@ -217,15 +323,47 @@ async def _check_access(
|
|||||||
if project_id is not None:
|
if project_id is not None:
|
||||||
project = await session.get(Project, project_id)
|
project = await session.get(Project, project_id)
|
||||||
if project is None:
|
if project is None:
|
||||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Project not found")
|
raise HTTPException(
|
||||||
|
status_code=status.HTTP_404_NOT_FOUND, detail="Project not found"
|
||||||
|
)
|
||||||
# Add ownership check if needed; for now just verify existence
|
# Add ownership check if needed; for now just verify existence
|
||||||
if tool_type_id is not None:
|
if tool_type_id is not None:
|
||||||
tool_type = await session.get(ToolType, tool_type_id)
|
tool_type = await session.get(ToolType, tool_type_id)
|
||||||
if tool_type is None:
|
if tool_type is None:
|
||||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Tool type not found")
|
raise HTTPException(
|
||||||
|
status_code=status.HTTP_404_NOT_FOUND, detail="Tool type not found"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
def _profile_to_response(profile: ConfigProfile, includes: list[ConfigProfileInclude] | None = None) -> dict:
|
async def _validate_git_mounts(
|
||||||
|
session: AsyncSession,
|
||||||
|
user_id: uuid.UUID,
|
||||||
|
git_mounts: list[Any],
|
||||||
|
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 {
|
return {
|
||||||
"id": str(profile.id),
|
"id": str(profile.id),
|
||||||
"user_id": str(profile.user_id),
|
"user_id": str(profile.user_id),
|
||||||
@@ -236,6 +374,7 @@ def _profile_to_response(profile: ConfigProfile, includes: list[ConfigProfileInc
|
|||||||
"env_vars": profile.env_vars or {},
|
"env_vars": profile.env_vars or {},
|
||||||
"runtime_hints": profile.runtime_hints or {},
|
"runtime_hints": profile.runtime_hints or {},
|
||||||
"mounts": profile.mounts or [],
|
"mounts": profile.mounts or [],
|
||||||
|
"git_mounts": profile.git_mounts or [],
|
||||||
"files": profile.files or {},
|
"files": profile.files or {},
|
||||||
"is_default": profile.is_default,
|
"is_default": profile.is_default,
|
||||||
"includes": [
|
"includes": [
|
||||||
@@ -254,13 +393,19 @@ def _profile_to_response(profile: ConfigProfile, includes: list[ConfigProfileInc
|
|||||||
@router.get("", response_model=list[ConfigProfileResponse])
|
@router.get("", response_model=list[ConfigProfileResponse])
|
||||||
async def list_config_profiles(
|
async def list_config_profiles(
|
||||||
project_id: str | None = Query(None, description="Filter by project compatibility"),
|
project_id: str | None = Query(None, description="Filter by project compatibility"),
|
||||||
tool_type_id: str | None = Query(None, description="Filter by tool type compatibility"),
|
tool_type_id: str | None = Query(
|
||||||
|
None, description="Filter by tool type compatibility"
|
||||||
|
),
|
||||||
current_user_id: uuid.UUID = Depends(get_current_user_id),
|
current_user_id: uuid.UUID = Depends(get_current_user_id),
|
||||||
session: AsyncSession = Depends(get_db_session),
|
session: AsyncSession = Depends(get_db_session),
|
||||||
):
|
):
|
||||||
"""List config profiles, optionally filtered by compatibility."""
|
"""List config profiles, optionally filtered by compatibility."""
|
||||||
user_uuid = current_user_id
|
user_uuid = current_user_id
|
||||||
query = select(ConfigProfile).where(ConfigProfile.user_id == user_uuid).options(selectinload(ConfigProfile.includes))
|
query = (
|
||||||
|
select(ConfigProfile)
|
||||||
|
.where(ConfigProfile.user_id == user_uuid)
|
||||||
|
.options(selectinload(ConfigProfile.includes))
|
||||||
|
)
|
||||||
|
|
||||||
if project_id or tool_type_id:
|
if project_id or tool_type_id:
|
||||||
# Compatibility filter: include portable profiles and matching scoped profiles
|
# Compatibility filter: include portable profiles and matching scoped profiles
|
||||||
@@ -272,7 +417,8 @@ async def list_config_profiles(
|
|||||||
conditions: list = []
|
conditions: list = []
|
||||||
# Portable profiles (no project, no tool)
|
# Portable profiles (no project, no tool)
|
||||||
conditions.append(
|
conditions.append(
|
||||||
(ConfigProfile.project_id.is_(None)) & (ConfigProfile.tool_type_id.is_(None))
|
(ConfigProfile.project_id.is_(None))
|
||||||
|
& (ConfigProfile.tool_type_id.is_(None))
|
||||||
)
|
)
|
||||||
if project_uuid:
|
if project_uuid:
|
||||||
# Profiles matching this project (with or without tool)
|
# Profiles matching this project (with or without tool)
|
||||||
@@ -283,7 +429,8 @@ async def list_config_profiles(
|
|||||||
if project_uuid and tool_uuid:
|
if project_uuid and tool_uuid:
|
||||||
# Exact match
|
# Exact match
|
||||||
conditions.append(
|
conditions.append(
|
||||||
(ConfigProfile.project_id == project_uuid) & (ConfigProfile.tool_type_id == tool_uuid)
|
(ConfigProfile.project_id == project_uuid)
|
||||||
|
& (ConfigProfile.tool_type_id == tool_uuid)
|
||||||
)
|
)
|
||||||
|
|
||||||
query = query.where(or_(*conditions))
|
query = query.where(or_(*conditions))
|
||||||
@@ -293,7 +440,9 @@ async def list_config_profiles(
|
|||||||
return [_profile_to_response(p) for p in profiles]
|
return [_profile_to_response(p) for p in profiles]
|
||||||
|
|
||||||
|
|
||||||
@router.post("", response_model=ConfigProfileResponse, status_code=status.HTTP_201_CREATED)
|
@router.post(
|
||||||
|
"", response_model=ConfigProfileResponse, status_code=status.HTTP_201_CREATED
|
||||||
|
)
|
||||||
async def create_config_profile(
|
async def create_config_profile(
|
||||||
data: ConfigProfileCreate,
|
data: ConfigProfileCreate,
|
||||||
current_user_id: uuid.UUID = Depends(get_current_user_id),
|
current_user_id: uuid.UUID = Depends(get_current_user_id),
|
||||||
@@ -304,10 +453,12 @@ async def create_config_profile(
|
|||||||
|
|
||||||
# Check for duplicate name
|
# Check for duplicate name
|
||||||
existing = await session.execute(
|
existing = await session.execute(
|
||||||
select(ConfigProfile).where(
|
select(ConfigProfile)
|
||||||
|
.where(
|
||||||
ConfigProfile.user_id == user_uuid,
|
ConfigProfile.user_id == user_uuid,
|
||||||
ConfigProfile.name == data.name,
|
ConfigProfile.name == data.name,
|
||||||
).options(selectinload(ConfigProfile.includes))
|
)
|
||||||
|
.options(selectinload(ConfigProfile.includes))
|
||||||
)
|
)
|
||||||
if existing.scalar_one_or_none() is not None:
|
if existing.scalar_one_or_none() is not None:
|
||||||
raise HTTPException(
|
raise HTTPException(
|
||||||
@@ -320,6 +471,13 @@ async def create_config_profile(
|
|||||||
tool_uuid = uuid.UUID(data.tool_type_id) if data.tool_type_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)
|
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
|
# Check size
|
||||||
size = _calculate_profile_size(data.model_dump())
|
size = _calculate_profile_size(data.model_dump())
|
||||||
if size > MAX_PROFILE_SIZE_BYTES:
|
if size > MAX_PROFILE_SIZE_BYTES:
|
||||||
@@ -337,6 +495,7 @@ async def create_config_profile(
|
|||||||
env_vars=data.env_vars,
|
env_vars=data.env_vars,
|
||||||
runtime_hints=data.runtime_hints,
|
runtime_hints=data.runtime_hints,
|
||||||
mounts=[m.model_dump() for m in data.mounts],
|
mounts=[m.model_dump() for m in data.mounts],
|
||||||
|
git_mounts=[m.model_dump() for m in data.git_mounts],
|
||||||
files=data.files,
|
files=data.files,
|
||||||
is_default=data.is_default,
|
is_default=data.is_default,
|
||||||
)
|
)
|
||||||
@@ -351,7 +510,7 @@ async def create_config_profile(
|
|||||||
)
|
)
|
||||||
profile = result.scalar_one()
|
profile = result.scalar_one()
|
||||||
|
|
||||||
logger.info("Created config profile %s for user %s", profile.id, user_uuid)
|
logger.debug("Created config profile %s for user %s", profile.id, user_uuid)
|
||||||
return _profile_to_response(profile)
|
return _profile_to_response(profile)
|
||||||
|
|
||||||
|
|
||||||
@@ -364,9 +523,13 @@ async def get_config_profile(
|
|||||||
"""Get a config profile by ID."""
|
"""Get a config profile by ID."""
|
||||||
profile = await _get_profile_with_includes(session, uuid.UUID(profile_id))
|
profile = await _get_profile_with_includes(session, uuid.UUID(profile_id))
|
||||||
if profile is None:
|
if profile is None:
|
||||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Profile not found")
|
raise HTTPException(
|
||||||
|
status_code=status.HTTP_404_NOT_FOUND, detail="Profile not found"
|
||||||
|
)
|
||||||
if profile.user_id != current_user_id:
|
if profile.user_id != current_user_id:
|
||||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Not authorized")
|
raise HTTPException(
|
||||||
|
status_code=status.HTTP_403_FORBIDDEN, detail="Not authorized"
|
||||||
|
)
|
||||||
return _profile_to_response(profile)
|
return _profile_to_response(profile)
|
||||||
|
|
||||||
|
|
||||||
@@ -380,9 +543,13 @@ async def update_config_profile(
|
|||||||
"""Update a config profile."""
|
"""Update a config profile."""
|
||||||
profile = await _get_profile_with_includes(session, uuid.UUID(profile_id))
|
profile = await _get_profile_with_includes(session, uuid.UUID(profile_id))
|
||||||
if profile is None:
|
if profile is None:
|
||||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Profile not found")
|
raise HTTPException(
|
||||||
|
status_code=status.HTTP_404_NOT_FOUND, detail="Profile not found"
|
||||||
|
)
|
||||||
if profile.user_id != current_user_id:
|
if profile.user_id != current_user_id:
|
||||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Not authorized")
|
raise HTTPException(
|
||||||
|
status_code=status.HTTP_403_FORBIDDEN, detail="Not authorized"
|
||||||
|
)
|
||||||
|
|
||||||
update_data = data.model_dump(exclude_unset=True)
|
update_data = data.model_dump(exclude_unset=True)
|
||||||
|
|
||||||
@@ -414,6 +581,16 @@ async def update_config_profile(
|
|||||||
)
|
)
|
||||||
await _check_access(session, profile.user_id, project_uuid, tool_uuid)
|
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
|
# Check size
|
||||||
current_data = _profile_to_response(profile)
|
current_data = _profile_to_response(profile)
|
||||||
merged = {**current_data, **update_data}
|
merged = {**current_data, **update_data}
|
||||||
@@ -430,6 +607,8 @@ async def update_config_profile(
|
|||||||
value = uuid.UUID(value) if value else None
|
value = uuid.UUID(value) if value else None
|
||||||
elif field_name == "mounts" and value is not 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]
|
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)
|
setattr(profile, field_name, value)
|
||||||
|
|
||||||
await session.commit()
|
await session.commit()
|
||||||
@@ -442,7 +621,7 @@ async def update_config_profile(
|
|||||||
)
|
)
|
||||||
profile = result.scalar_one()
|
profile = result.scalar_one()
|
||||||
|
|
||||||
logger.info("Updated config profile %s", profile.id)
|
logger.debug("Updated config profile %s", profile.id)
|
||||||
return _profile_to_response(profile)
|
return _profile_to_response(profile)
|
||||||
|
|
||||||
|
|
||||||
@@ -455,14 +634,18 @@ async def delete_config_profile(
|
|||||||
"""Delete a config profile."""
|
"""Delete a config profile."""
|
||||||
profile = await _get_profile_with_includes(session, uuid.UUID(profile_id))
|
profile = await _get_profile_with_includes(session, uuid.UUID(profile_id))
|
||||||
if profile is None:
|
if profile is None:
|
||||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Profile not found")
|
raise HTTPException(
|
||||||
|
status_code=status.HTTP_404_NOT_FOUND, detail="Profile not found"
|
||||||
|
)
|
||||||
if profile.user_id != current_user_id:
|
if profile.user_id != current_user_id:
|
||||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Not authorized")
|
raise HTTPException(
|
||||||
|
status_code=status.HTTP_403_FORBIDDEN, detail="Not authorized"
|
||||||
|
)
|
||||||
|
|
||||||
await session.delete(profile)
|
await session.delete(profile)
|
||||||
await session.commit()
|
await session.commit()
|
||||||
|
|
||||||
logger.info("Deleted config profile %s", profile_id)
|
logger.debug("Deleted config profile %s", profile_id)
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
|
||||||
@@ -476,9 +659,13 @@ async def update_profile_includes(
|
|||||||
"""Update the ordered includes for a config profile."""
|
"""Update the ordered includes for a config profile."""
|
||||||
profile = await _get_profile_with_includes(session, uuid.UUID(profile_id))
|
profile = await _get_profile_with_includes(session, uuid.UUID(profile_id))
|
||||||
if profile is None:
|
if profile is None:
|
||||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Profile not found")
|
raise HTTPException(
|
||||||
|
status_code=status.HTTP_404_NOT_FOUND, detail="Profile not found"
|
||||||
|
)
|
||||||
if profile.user_id != current_user_id:
|
if profile.user_id != current_user_id:
|
||||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Not authorized")
|
raise HTTPException(
|
||||||
|
status_code=status.HTTP_403_FORBIDDEN, detail="Not authorized"
|
||||||
|
)
|
||||||
|
|
||||||
# Validate all included profiles exist and belong to the user
|
# Validate all included profiles exist and belong to the user
|
||||||
included_uuids = [uuid.UUID(inc_id) for inc_id in data.includes]
|
included_uuids = [uuid.UUID(inc_id) for inc_id in data.includes]
|
||||||
@@ -518,7 +705,9 @@ async def update_profile_includes(
|
|||||||
|
|
||||||
# Remove existing includes
|
# Remove existing includes
|
||||||
result = await session.execute(
|
result = await session.execute(
|
||||||
select(ConfigProfileInclude).where(ConfigProfileInclude.profile_id == profile.id)
|
select(ConfigProfileInclude).where(
|
||||||
|
ConfigProfileInclude.profile_id == profile.id
|
||||||
|
)
|
||||||
)
|
)
|
||||||
for existing in result.scalars().all():
|
for existing in result.scalars().all():
|
||||||
await session.delete(existing)
|
await session.delete(existing)
|
||||||
@@ -543,11 +732,13 @@ async def update_profile_includes(
|
|||||||
profile = result.scalar_one()
|
profile = result.scalar_one()
|
||||||
|
|
||||||
inc_result = await session.execute(
|
inc_result = await session.execute(
|
||||||
select(ConfigProfileInclude).where(ConfigProfileInclude.profile_id == profile.id)
|
select(ConfigProfileInclude).where(
|
||||||
|
ConfigProfileInclude.profile_id == profile.id
|
||||||
|
)
|
||||||
)
|
)
|
||||||
direct_includes = inc_result.scalars().all()
|
direct_includes = inc_result.scalars().all()
|
||||||
|
|
||||||
logger.info("Updated includes for config profile %s", profile.id)
|
logger.debug("Updated includes for config profile %s", profile.id)
|
||||||
return _profile_to_response(profile, list(direct_includes))
|
return _profile_to_response(profile, list(direct_includes))
|
||||||
|
|
||||||
|
|
||||||
@@ -560,9 +751,13 @@ async def preview_config_profile(
|
|||||||
"""Preview the resolved output of a config profile."""
|
"""Preview the resolved output of a config profile."""
|
||||||
profile = await _get_profile_with_includes(session, uuid.UUID(profile_id))
|
profile = await _get_profile_with_includes(session, uuid.UUID(profile_id))
|
||||||
if profile is None:
|
if profile is None:
|
||||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Profile not found")
|
raise HTTPException(
|
||||||
|
status_code=status.HTTP_404_NOT_FOUND, detail="Profile not found"
|
||||||
|
)
|
||||||
if profile.user_id != current_user_id:
|
if profile.user_id != current_user_id:
|
||||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Not authorized")
|
raise HTTPException(
|
||||||
|
status_code=status.HTTP_403_FORBIDDEN, detail="Not authorized"
|
||||||
|
)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
resolved = await resolve_profile(session, profile.id)
|
resolved = await resolve_profile(session, profile.id)
|
||||||
@@ -643,3 +838,171 @@ async def resolve_default_profile(
|
|||||||
# Fall back to first created compatible profile
|
# Fall back to first created compatible profile
|
||||||
first = profiles[0]
|
first = profiles[0]
|
||||||
return {"profile_id": str(first.id), "profile_name": first.name}
|
return {"profile_id": str(first.id), "profile_name": first.name}
|
||||||
|
|
||||||
|
|
||||||
|
class ValidateGitUrlRequest(BaseModel):
|
||||||
|
url: str = Field(description="Git remote URL to validate")
|
||||||
|
ssh_key_id: str | None = Field(
|
||||||
|
default=None, description="Optional SSH key ID for private repos"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class ValidateGitUrlResponse(BaseModel):
|
||||||
|
valid: bool
|
||||||
|
suggested_url: str | None = None
|
||||||
|
branches: list[str] | None = None
|
||||||
|
default_branch: str | None = None
|
||||||
|
error: str | None = None
|
||||||
|
error_code: str | None = None
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("/validate-git-url", response_model=ValidateGitUrlResponse)
|
||||||
|
async def validate_git_url(
|
||||||
|
data: ValidateGitUrlRequest,
|
||||||
|
current_user_id: uuid.UUID = Depends(get_current_user_id),
|
||||||
|
session: AsyncSession = Depends(get_db_session),
|
||||||
|
) -> ValidateGitUrlResponse:
|
||||||
|
"""Validate a git remote URL and list available branches.
|
||||||
|
|
||||||
|
Parses the URL, suggests corrections for browser URLs, and runs
|
||||||
|
git ls-remote to verify reachability and enumerate branches.
|
||||||
|
"""
|
||||||
|
parse_result = parse_git_url(data.url)
|
||||||
|
original_url = data.url.strip()
|
||||||
|
url_to_check = parse_result.get("base_url") or original_url
|
||||||
|
|
||||||
|
if not url_to_check:
|
||||||
|
return ValidateGitUrlResponse(
|
||||||
|
valid=False,
|
||||||
|
error=parse_result.get("message", "Invalid URL"),
|
||||||
|
error_code=parse_result.get("error_code", "INVALID_URL"),
|
||||||
|
)
|
||||||
|
|
||||||
|
# If the URL needed parsing, return suggestion without checking remote
|
||||||
|
if parse_result.get("needs_parsing") and url_to_check != original_url:
|
||||||
|
return ValidateGitUrlResponse(
|
||||||
|
valid=False,
|
||||||
|
suggested_url=url_to_check,
|
||||||
|
error=parse_result.get("message"),
|
||||||
|
error_code=parse_result.get("error_code", "URL_NEEDS_PARSING"),
|
||||||
|
)
|
||||||
|
|
||||||
|
# Optional SSH key for private repos
|
||||||
|
env = None
|
||||||
|
key_path = None
|
||||||
|
if data.ssh_key_id:
|
||||||
|
from src.models.ssh_key import SSHKey
|
||||||
|
from src.services.ssh_keys import _get_fernet
|
||||||
|
|
||||||
|
try:
|
||||||
|
ssh_key_uuid = uuid.UUID(data.ssh_key_id)
|
||||||
|
except ValueError:
|
||||||
|
return ValidateGitUrlResponse(
|
||||||
|
valid=False,
|
||||||
|
error="Invalid SSH key ID format",
|
||||||
|
error_code="INVALID_SSH_KEY",
|
||||||
|
)
|
||||||
|
|
||||||
|
ssh_key = await session.get(SSHKey, ssh_key_uuid)
|
||||||
|
if ssh_key is None or ssh_key.user_id != current_user_id:
|
||||||
|
return ValidateGitUrlResponse(
|
||||||
|
valid=False,
|
||||||
|
error="SSH key not found or not authorized",
|
||||||
|
error_code="SSH_KEY_NOT_FOUND",
|
||||||
|
)
|
||||||
|
|
||||||
|
import tempfile
|
||||||
|
|
||||||
|
fernet = _get_fernet()
|
||||||
|
private_key = fernet.decrypt(ssh_key.private_key_encrypted.encode()).decode()
|
||||||
|
fd, key_path = tempfile.mkstemp(prefix="ssh_key_")
|
||||||
|
try:
|
||||||
|
os.write(fd, private_key.encode())
|
||||||
|
finally:
|
||||||
|
os.close(fd)
|
||||||
|
os.chmod(key_path, 0o600)
|
||||||
|
env = {
|
||||||
|
"GIT_SSH_COMMAND": f"ssh -i {key_path} -o StrictHostKeyChecking=no -o UserKnownHostsFile=/dev/null"
|
||||||
|
}
|
||||||
|
|
||||||
|
try:
|
||||||
|
result = subprocess.run(
|
||||||
|
["git", "ls-remote", "--heads", url_to_check],
|
||||||
|
capture_output=True,
|
||||||
|
text=True,
|
||||||
|
timeout=30,
|
||||||
|
env={**os.environ, **env} if env else None,
|
||||||
|
)
|
||||||
|
except subprocess.TimeoutExpired:
|
||||||
|
if key_path and os.path.exists(key_path):
|
||||||
|
os.unlink(key_path)
|
||||||
|
return ValidateGitUrlResponse(
|
||||||
|
valid=False,
|
||||||
|
error="Remote repository check timed out",
|
||||||
|
error_code="TIMEOUT",
|
||||||
|
)
|
||||||
|
except FileNotFoundError:
|
||||||
|
if key_path and os.path.exists(key_path):
|
||||||
|
os.unlink(key_path)
|
||||||
|
return ValidateGitUrlResponse(
|
||||||
|
valid=False,
|
||||||
|
error="git command not found on server",
|
||||||
|
error_code="GIT_NOT_FOUND",
|
||||||
|
)
|
||||||
|
finally:
|
||||||
|
if key_path and os.path.exists(key_path):
|
||||||
|
os.unlink(key_path)
|
||||||
|
|
||||||
|
if result.returncode != 0:
|
||||||
|
stderr = result.stderr.strip()
|
||||||
|
if (
|
||||||
|
"could not resolve" in stderr.lower()
|
||||||
|
or "unable to access" in stderr.lower()
|
||||||
|
):
|
||||||
|
error_msg = "Could not reach repository. Check the URL and network access."
|
||||||
|
error_code = "UNREACHABLE"
|
||||||
|
elif (
|
||||||
|
"authentication" in stderr.lower() or "permission denied" in stderr.lower()
|
||||||
|
):
|
||||||
|
error_msg = (
|
||||||
|
"Authentication failed. Provide an SSH key for private repositories."
|
||||||
|
)
|
||||||
|
error_code = "AUTH_FAILED"
|
||||||
|
else:
|
||||||
|
error_msg = f"Repository not accessible: {stderr[:200]}"
|
||||||
|
error_code = "REMOTE_ERROR"
|
||||||
|
return ValidateGitUrlResponse(
|
||||||
|
valid=False,
|
||||||
|
error=error_msg,
|
||||||
|
error_code=error_code,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Parse branches from ls-remote output
|
||||||
|
branches: list[str] = []
|
||||||
|
default_branch = "main"
|
||||||
|
for line in result.stdout.strip().split("\n"):
|
||||||
|
if not line.strip():
|
||||||
|
continue
|
||||||
|
parts = line.split()
|
||||||
|
if len(parts) == 2:
|
||||||
|
ref = parts[1]
|
||||||
|
# refs/heads/branch-name
|
||||||
|
if ref.startswith("refs/heads/"):
|
||||||
|
branch_name = ref[len("refs/heads/") :]
|
||||||
|
branches.append(branch_name)
|
||||||
|
if branch_name in ("main", "master"):
|
||||||
|
default_branch = branch_name
|
||||||
|
|
||||||
|
if not branches:
|
||||||
|
return ValidateGitUrlResponse(
|
||||||
|
valid=False,
|
||||||
|
error="No branches found in remote repository",
|
||||||
|
error_code="NO_BRANCHES",
|
||||||
|
)
|
||||||
|
|
||||||
|
return ValidateGitUrlResponse(
|
||||||
|
valid=True,
|
||||||
|
suggested_url=url_to_check if url_to_check != original_url else None,
|
||||||
|
branches=branches,
|
||||||
|
default_branch=default_branch,
|
||||||
|
)
|
||||||
|
|||||||
@@ -0,0 +1,80 @@
|
|||||||
|
"""SSE streaming endpoint for instance events."""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import contextlib
|
||||||
|
import json
|
||||||
|
import uuid
|
||||||
|
from collections.abc import AsyncGenerator
|
||||||
|
|
||||||
|
from fastapi import APIRouter, Depends, HTTPException, Request, status
|
||||||
|
from fastapi.responses import StreamingResponse
|
||||||
|
|
||||||
|
from src.auth.dependencies import get_current_user_id
|
||||||
|
from src.services.event_bus import InstanceEventBus, InstanceEventPayload
|
||||||
|
|
||||||
|
router = APIRouter(prefix="/events", tags=["events"])
|
||||||
|
|
||||||
|
# In-memory connection counter per user (single-process assumption)
|
||||||
|
_connection_counts: dict[uuid.UUID, int] = {}
|
||||||
|
MAX_CONNECTIONS_PER_USER = 5
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/stream")
|
||||||
|
async def events_stream(
|
||||||
|
request: Request,
|
||||||
|
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||||
|
) -> StreamingResponse:
|
||||||
|
"""Stream instance events via Server-Sent Events.
|
||||||
|
|
||||||
|
Enforces a maximum of 5 concurrent connections per user.
|
||||||
|
"""
|
||||||
|
current = _connection_counts.get(user_id, 0)
|
||||||
|
if current >= MAX_CONNECTIONS_PER_USER:
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=status.HTTP_429_TOO_MANY_REQUESTS,
|
||||||
|
detail="Too many SSE connections",
|
||||||
|
)
|
||||||
|
|
||||||
|
_connection_counts[user_id] = current + 1
|
||||||
|
|
||||||
|
async def event_generator() -> AsyncGenerator[str, None]:
|
||||||
|
event_bus = InstanceEventBus()
|
||||||
|
queue: asyncio.Queue[InstanceEventPayload] = asyncio.Queue(maxsize=100)
|
||||||
|
|
||||||
|
async def on_event(payload: InstanceEventPayload) -> None:
|
||||||
|
try:
|
||||||
|
queue.put_nowait(payload)
|
||||||
|
except asyncio.QueueFull:
|
||||||
|
# Drop oldest event to make room
|
||||||
|
with contextlib.suppress(asyncio.QueueEmpty):
|
||||||
|
queue.get_nowait()
|
||||||
|
with contextlib.suppress(asyncio.QueueFull):
|
||||||
|
queue.put_nowait(payload)
|
||||||
|
|
||||||
|
unsubscribe = event_bus.subscribe("*", on_event)
|
||||||
|
|
||||||
|
try:
|
||||||
|
while True:
|
||||||
|
try:
|
||||||
|
payload = await asyncio.wait_for(queue.get(), timeout=30.0)
|
||||||
|
yield f"event: {payload['event']}\ndata: {json.dumps(payload)}\n\n"
|
||||||
|
except asyncio.TimeoutError:
|
||||||
|
yield ":ping\n\n"
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
# Client disconnected
|
||||||
|
raise
|
||||||
|
finally:
|
||||||
|
unsubscribe()
|
||||||
|
_connection_counts[user_id] = max(0, _connection_counts.get(user_id, 1) - 1)
|
||||||
|
if _connection_counts[user_id] == 0:
|
||||||
|
_connection_counts.pop(user_id, None)
|
||||||
|
|
||||||
|
return StreamingResponse(
|
||||||
|
event_generator(),
|
||||||
|
media_type="text/event-stream",
|
||||||
|
headers={
|
||||||
|
"Cache-Control": "no-cache",
|
||||||
|
"Connection": "keep-alive",
|
||||||
|
"X-Accel-Buffering": "no",
|
||||||
|
},
|
||||||
|
)
|
||||||
@@ -10,12 +10,10 @@ from pydantic import BaseModel, ConfigDict
|
|||||||
from sqlalchemy import select
|
from sqlalchemy import select
|
||||||
from sqlalchemy.ext.asyncio import AsyncSession
|
from sqlalchemy.ext.asyncio import AsyncSession
|
||||||
|
|
||||||
from src.auth.dependencies import get_current_user_id, get_db_session
|
from src.auth.dependencies import _get_owned_project, _get_user, get_current_user_id, get_db_session
|
||||||
from src.config import Settings
|
from src.config import Settings
|
||||||
from src.models.git_repository import GitRepository
|
from src.models.git_repository import GitRepository
|
||||||
from src.models.project import Project
|
|
||||||
from src.models.ssh_key import SSHKey
|
from src.models.ssh_key import SSHKey
|
||||||
from src.models.user import User
|
|
||||||
from src.utils.git_files import (
|
from src.utils.git_files import (
|
||||||
commit_file,
|
commit_file,
|
||||||
get_file_content,
|
get_file_content,
|
||||||
@@ -42,40 +40,6 @@ router = APIRouter(prefix="/projects", tags=["git-repositories"])
|
|||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
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: If project not found or user is not the owner.
|
|
||||||
"""
|
|
||||||
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
|
|
||||||
|
|
||||||
|
|
||||||
def _get_repo_path(user_id: uuid.UUID, project_id: uuid.UUID, name: str) -> str:
|
def _get_repo_path(user_id: uuid.UUID, project_id: uuid.UUID, name: str) -> str:
|
||||||
"""Generate the filesystem path for a repository.
|
"""Generate the filesystem path for a repository.
|
||||||
|
|
||||||
@@ -256,7 +220,7 @@ class GitRepositoryResponse(BaseModel):
|
|||||||
id: uuid.UUID
|
id: uuid.UUID
|
||||||
name: str
|
name: str
|
||||||
path: str
|
path: str
|
||||||
project_id: uuid.UUID
|
project_id: uuid.UUID | None
|
||||||
owner_id: uuid.UUID
|
owner_id: uuid.UUID
|
||||||
is_mirror: bool
|
is_mirror: bool
|
||||||
remote_url: str | None
|
remote_url: str | None
|
||||||
@@ -266,6 +230,156 @@ class GitRepositoryResponse(BaseModel):
|
|||||||
updated_at: datetime
|
updated_at: datetime
|
||||||
|
|
||||||
|
|
||||||
|
@router.get(
|
||||||
|
"/repositories",
|
||||||
|
response_model=list[GitRepositoryResponse],
|
||||||
|
summary="List all user repositories",
|
||||||
|
description="List all git repositories owned by the user, including external repositories not tied to any project.",
|
||||||
|
)
|
||||||
|
async def list_user_repositories(
|
||||||
|
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||||
|
session: AsyncSession = Depends(get_db_session),
|
||||||
|
) -> list[GitRepository]:
|
||||||
|
"""List all repositories owned by the user.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
user_id: ID of the authenticated user.
|
||||||
|
session: Database session.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
List of all repositories owned by the user.
|
||||||
|
"""
|
||||||
|
result = await session.execute(
|
||||||
|
select(GitRepository).where(GitRepository.owner_id == user_id)
|
||||||
|
)
|
||||||
|
return list(result.scalars().all())
|
||||||
|
|
||||||
|
|
||||||
|
@router.post(
|
||||||
|
"/repositories/parse-url",
|
||||||
|
response_model=URLParseResponse,
|
||||||
|
summary="Parse a git URL",
|
||||||
|
description="Parse a git URL and detect if it's a browser URL that needs correction.",
|
||||||
|
)
|
||||||
|
async def parse_repository_url(data: URLParseRequest) -> URLParseResponse:
|
||||||
|
"""Parse a git URL and detect if it's a browser URL that needs correction.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
data: Request containing the URL to parse.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Parsed URL information including whether it needs parsing and suggested corrections.
|
||||||
|
"""
|
||||||
|
result = parse_git_url(data.url)
|
||||||
|
return URLParseResponse(**result)
|
||||||
|
|
||||||
|
|
||||||
|
@router.post(
|
||||||
|
"/repositories",
|
||||||
|
response_model=GitRepositoryResponse,
|
||||||
|
status_code=status.HTTP_201_CREATED,
|
||||||
|
summary="Create an external repository",
|
||||||
|
description="Create a new external git repository (not tied to any project). Can clone from remote URL.",
|
||||||
|
)
|
||||||
|
async def create_external_repository(
|
||||||
|
data: GitRepositoryCreate,
|
||||||
|
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||||
|
session: AsyncSession = Depends(get_db_session),
|
||||||
|
) -> GitRepository:
|
||||||
|
"""Create a new external git repository.
|
||||||
|
|
||||||
|
External repositories are not tied to any project and can be used
|
||||||
|
across all projects for config profile git mounts.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
data: Repository creation data including name and optional remote URL.
|
||||||
|
user_id: ID of the authenticated user.
|
||||||
|
session: Database session.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
The newly created external repository.
|
||||||
|
"""
|
||||||
|
_user = await _get_user(session, user_id)
|
||||||
|
|
||||||
|
# Check for duplicate name (external repos only)
|
||||||
|
existing = await session.execute(
|
||||||
|
select(GitRepository).where(
|
||||||
|
GitRepository.project_id.is_(None),
|
||||||
|
GitRepository.owner_id == user_id,
|
||||||
|
GitRepository.name == data.name,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
if existing.scalar_one_or_none():
|
||||||
|
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="repository name already exists")
|
||||||
|
|
||||||
|
# Validate and potentially correct the URL
|
||||||
|
remote_url = data.remote_url
|
||||||
|
if remote_url and not data.force_original_url:
|
||||||
|
parse_result = parse_git_url(remote_url)
|
||||||
|
if parse_result["needs_parsing"] and parse_result["base_url"]:
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
|
||||||
|
detail={
|
||||||
|
"message": "The provided URL appears to be a browser URL, not a git clone URL",
|
||||||
|
"suggested_url": parse_result["base_url"],
|
||||||
|
"original_url": remote_url,
|
||||||
|
"error_code": "URL_NEEDS_PARSING",
|
||||||
|
},
|
||||||
|
)
|
||||||
|
if parse_result["base_url"]:
|
||||||
|
remote_url = parse_result["base_url"]
|
||||||
|
|
||||||
|
# Validate SSH key if provided
|
||||||
|
ssh_key_id = None
|
||||||
|
ssh_key = None
|
||||||
|
if data.ssh_key_id:
|
||||||
|
try:
|
||||||
|
ssh_key_id = uuid.UUID(data.ssh_key_id)
|
||||||
|
except ValueError:
|
||||||
|
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="invalid ssh_key_id format")
|
||||||
|
|
||||||
|
ssh_key = await session.get(SSHKey, ssh_key_id)
|
||||||
|
if ssh_key is None:
|
||||||
|
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="ssh key not found")
|
||||||
|
if ssh_key.user_id != user_id:
|
||||||
|
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="ssh key does not belong to user")
|
||||||
|
|
||||||
|
if remote_url:
|
||||||
|
_preflight_remote_repository(remote_url, ssh_key)
|
||||||
|
|
||||||
|
# Create external repo with no project
|
||||||
|
repo = GitRepository(
|
||||||
|
name=data.name,
|
||||||
|
path="", # Will be set after clone
|
||||||
|
project_id=None,
|
||||||
|
owner_id=user_id,
|
||||||
|
remote_url=remote_url,
|
||||||
|
ssh_key_id=ssh_key_id,
|
||||||
|
)
|
||||||
|
session.add(repo)
|
||||||
|
await session.flush()
|
||||||
|
|
||||||
|
# Set path and optionally clone
|
||||||
|
repo_path = f"/data/repos/external/{user_id}/{repo.id}"
|
||||||
|
repo.path = repo_path
|
||||||
|
|
||||||
|
if remote_url:
|
||||||
|
try:
|
||||||
|
_clone_working_repository(remote_url, repo_path, ssh_key)
|
||||||
|
repo.is_mirror = False
|
||||||
|
except Exception as exc:
|
||||||
|
await session.rollback()
|
||||||
|
raise HTTPException(status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail=f"Failed to clone repository: {exc}")
|
||||||
|
else:
|
||||||
|
# Initialize empty repo
|
||||||
|
os.makedirs(repo_path, exist_ok=True)
|
||||||
|
subprocess.run(["git", "init", repo_path], check=True, capture_output=True)
|
||||||
|
repo.is_mirror = False
|
||||||
|
|
||||||
|
await session.commit()
|
||||||
|
return repo
|
||||||
|
|
||||||
|
|
||||||
@router.get(
|
@router.get(
|
||||||
"/{project_id}/repositories",
|
"/{project_id}/repositories",
|
||||||
response_model=list[GitRepositoryResponse],
|
response_model=list[GitRepositoryResponse],
|
||||||
@@ -335,25 +449,6 @@ async def delete_repository(
|
|||||||
return Response(status_code=status.HTTP_204_NO_CONTENT)
|
return Response(status_code=status.HTTP_204_NO_CONTENT)
|
||||||
|
|
||||||
|
|
||||||
@router.post(
|
|
||||||
"/repositories/parse-url",
|
|
||||||
response_model=URLParseResponse,
|
|
||||||
summary="Parse a git URL",
|
|
||||||
description="Parse a git URL and detect if it's a browser URL that needs correction.",
|
|
||||||
)
|
|
||||||
async def parse_repository_url(data: URLParseRequest) -> URLParseResponse:
|
|
||||||
"""Parse a git URL and detect if it's a browser URL that needs correction.
|
|
||||||
|
|
||||||
Args:
|
|
||||||
data: Request containing the URL to parse.
|
|
||||||
|
|
||||||
Returns:
|
|
||||||
Parsed URL information including whether it needs parsing and suggested corrections.
|
|
||||||
"""
|
|
||||||
result = parse_git_url(data.url)
|
|
||||||
return URLParseResponse(**result)
|
|
||||||
|
|
||||||
|
|
||||||
@router.post(
|
@router.post(
|
||||||
"/{project_id}/repositories",
|
"/{project_id}/repositories",
|
||||||
response_model=GitRepositoryResponse,
|
response_model=GitRepositoryResponse,
|
||||||
|
|||||||
@@ -4,11 +4,10 @@ import time
|
|||||||
from datetime import datetime, timezone
|
from datetime import datetime, timezone
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
from fastapi import APIRouter, status
|
from fastapi import APIRouter
|
||||||
from pydantic import BaseModel, Field
|
from pydantic import BaseModel, Field
|
||||||
from sqlalchemy import text
|
from sqlalchemy import text
|
||||||
|
|
||||||
from src.config import Settings
|
|
||||||
from src.database import SessionLocal
|
from src.database import SessionLocal
|
||||||
|
|
||||||
router = APIRouter()
|
router = APIRouter()
|
||||||
|
|||||||
@@ -2,7 +2,6 @@
|
|||||||
|
|
||||||
import logging
|
import logging
|
||||||
import uuid
|
import uuid
|
||||||
from typing import Any
|
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
from fastapi import APIRouter, Depends, HTTPException, Request, Response, status
|
from fastapi import APIRouter, Depends, HTTPException, Request, Response, status
|
||||||
|
|||||||
@@ -0,0 +1,161 @@
|
|||||||
|
"""Notification API endpoints."""
|
||||||
|
|
||||||
|
import uuid
|
||||||
|
from datetime import datetime
|
||||||
|
|
||||||
|
from fastapi import APIRouter, Depends, HTTPException, Query, status
|
||||||
|
from pydantic import BaseModel, ConfigDict, Field
|
||||||
|
from sqlalchemy.ext.asyncio import AsyncSession
|
||||||
|
|
||||||
|
from src.auth.dependencies import get_current_user, get_db_session
|
||||||
|
from src.models.user import User
|
||||||
|
from src.models.user_config import UserConfig
|
||||||
|
from src.services.notification_service import notification_service
|
||||||
|
|
||||||
|
router = APIRouter(prefix="/notifications", tags=["notifications"])
|
||||||
|
|
||||||
|
|
||||||
|
class NotificationItem(BaseModel):
|
||||||
|
model_config = ConfigDict(from_attributes=True)
|
||||||
|
|
||||||
|
id: uuid.UUID
|
||||||
|
user_id: uuid.UUID
|
||||||
|
category: str
|
||||||
|
severity: str
|
||||||
|
title: str
|
||||||
|
message: str | None
|
||||||
|
source_type: str | None
|
||||||
|
source_id: uuid.UUID | None
|
||||||
|
notification_metadata: dict = Field(serialization_alias="metadata")
|
||||||
|
read_at: datetime | None
|
||||||
|
dismissed_at: datetime | None
|
||||||
|
created_at: datetime
|
||||||
|
|
||||||
|
|
||||||
|
class NotificationListResponse(BaseModel):
|
||||||
|
items: list[NotificationItem]
|
||||||
|
total: int
|
||||||
|
limit: int
|
||||||
|
offset: int
|
||||||
|
|
||||||
|
|
||||||
|
class UnreadCountResponse(BaseModel):
|
||||||
|
count: int
|
||||||
|
|
||||||
|
|
||||||
|
class MarkAllReadResponse(BaseModel):
|
||||||
|
marked_count: int
|
||||||
|
|
||||||
|
|
||||||
|
class ClearAllResponse(BaseModel):
|
||||||
|
cleared_count: int
|
||||||
|
|
||||||
|
|
||||||
|
async def _get_mute_categories(
|
||||||
|
session: AsyncSession,
|
||||||
|
user_id: uuid.UUID,
|
||||||
|
) -> list[str]:
|
||||||
|
"""Read notification mute categories from user config."""
|
||||||
|
from sqlalchemy import select
|
||||||
|
|
||||||
|
result = await session.execute(
|
||||||
|
select(UserConfig).where(UserConfig.user_id == user_id)
|
||||||
|
)
|
||||||
|
config = result.scalar_one_or_none()
|
||||||
|
if config is None:
|
||||||
|
return []
|
||||||
|
mute_categories = config.config.get("notification_mute_categories", [])
|
||||||
|
if isinstance(mute_categories, list):
|
||||||
|
return mute_categories
|
||||||
|
return []
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("", response_model=NotificationListResponse)
|
||||||
|
async def list_notifications(
|
||||||
|
limit: int = Query(20, ge=1, le=100),
|
||||||
|
offset: int = Query(0, ge=0),
|
||||||
|
unread_only: bool = Query(False),
|
||||||
|
user: User = Depends(get_current_user),
|
||||||
|
session: AsyncSession = Depends(get_db_session),
|
||||||
|
) -> NotificationListResponse:
|
||||||
|
"""List notifications for the authenticated user."""
|
||||||
|
mute_categories = await _get_mute_categories(session, user.id)
|
||||||
|
items, total = await notification_service.list_notifications(
|
||||||
|
session,
|
||||||
|
user.id,
|
||||||
|
limit=limit,
|
||||||
|
offset=offset,
|
||||||
|
unread_only=unread_only,
|
||||||
|
mute_categories=mute_categories,
|
||||||
|
)
|
||||||
|
return NotificationListResponse(
|
||||||
|
items=[NotificationItem.model_validate(item) for item in items],
|
||||||
|
total=total,
|
||||||
|
limit=limit,
|
||||||
|
offset=offset,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@router.get("/unread", response_model=UnreadCountResponse)
|
||||||
|
async def get_unread_count(
|
||||||
|
user: User = Depends(get_current_user),
|
||||||
|
session: AsyncSession = Depends(get_db_session),
|
||||||
|
) -> UnreadCountResponse:
|
||||||
|
"""Get unread notification count for the authenticated user."""
|
||||||
|
count = await notification_service.get_unread_count(session, user.id)
|
||||||
|
return UnreadCountResponse(count=count)
|
||||||
|
|
||||||
|
|
||||||
|
@router.patch("/{notification_id}/read", response_model=NotificationItem)
|
||||||
|
async def mark_notification_read(
|
||||||
|
notification_id: uuid.UUID,
|
||||||
|
user: User = Depends(get_current_user),
|
||||||
|
session: AsyncSession = Depends(get_db_session),
|
||||||
|
) -> NotificationItem:
|
||||||
|
"""Mark a single notification as read."""
|
||||||
|
try:
|
||||||
|
notification = await notification_service.mark_read(
|
||||||
|
session, notification_id, user.id
|
||||||
|
)
|
||||||
|
except ValueError as exc:
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=status.HTTP_404_NOT_FOUND,
|
||||||
|
detail="Notification not found",
|
||||||
|
) from exc
|
||||||
|
return NotificationItem.model_validate(notification)
|
||||||
|
|
||||||
|
|
||||||
|
@router.post("/mark-all-read", response_model=MarkAllReadResponse)
|
||||||
|
async def mark_all_read(
|
||||||
|
user: User = Depends(get_current_user),
|
||||||
|
session: AsyncSession = Depends(get_db_session),
|
||||||
|
) -> MarkAllReadResponse:
|
||||||
|
"""Mark all unread notifications as read."""
|
||||||
|
marked = await notification_service.mark_all_read(session, user.id)
|
||||||
|
return MarkAllReadResponse(marked_count=marked)
|
||||||
|
|
||||||
|
|
||||||
|
@router.delete("", status_code=status.HTTP_200_OK)
|
||||||
|
async def clear_all_notifications(
|
||||||
|
user: User = Depends(get_current_user),
|
||||||
|
session: AsyncSession = Depends(get_db_session),
|
||||||
|
) -> ClearAllResponse:
|
||||||
|
"""Dismiss all notifications for the authenticated user."""
|
||||||
|
cleared = await notification_service.dismiss_all(session, user.id)
|
||||||
|
return ClearAllResponse(cleared_count=cleared)
|
||||||
|
|
||||||
|
|
||||||
|
@router.delete("/{notification_id}", status_code=status.HTTP_204_NO_CONTENT)
|
||||||
|
async def dismiss_notification(
|
||||||
|
notification_id: uuid.UUID,
|
||||||
|
user: User = Depends(get_current_user),
|
||||||
|
session: AsyncSession = Depends(get_db_session),
|
||||||
|
) -> None:
|
||||||
|
"""Soft-delete (dismiss) a single notification."""
|
||||||
|
try:
|
||||||
|
await notification_service.dismiss(session, notification_id, user.id)
|
||||||
|
except ValueError as exc:
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=status.HTTP_404_NOT_FOUND,
|
||||||
|
detail="Notification not found",
|
||||||
|
) from exc
|
||||||
@@ -7,23 +7,14 @@ from pydantic import BaseModel, ConfigDict
|
|||||||
from sqlalchemy import select
|
from sqlalchemy import select
|
||||||
from sqlalchemy.ext.asyncio import AsyncSession
|
from sqlalchemy.ext.asyncio import AsyncSession
|
||||||
|
|
||||||
from src.auth.dependencies import get_current_user_id, get_db_session
|
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.git_repository import GitRepository
|
||||||
from src.models.project import Project
|
from src.models.project import Project
|
||||||
from src.models.ssh_key import SSHKey
|
from src.models.ssh_key import SSHKey
|
||||||
from src.models.user import User
|
|
||||||
|
|
||||||
router = APIRouter(prefix="/projects", tags=["projects"])
|
router = APIRouter(prefix="/projects", tags=["projects"])
|
||||||
|
|
||||||
|
|
||||||
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
|
|
||||||
|
|
||||||
|
|
||||||
class ProjectCreate(BaseModel):
|
class ProjectCreate(BaseModel):
|
||||||
name: str
|
name: str
|
||||||
description: str | None = None
|
description: str | None = None
|
||||||
@@ -132,32 +123,6 @@ async def get_project(
|
|||||||
return await _get_owned_project(project_id, user_id, session)
|
return await _get_owned_project(project_id, user_id, session)
|
||||||
|
|
||||||
|
|
||||||
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: If project not found or user is not the owner.
|
|
||||||
"""
|
|
||||||
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
|
|
||||||
|
|
||||||
|
|
||||||
@router.patch(
|
@router.patch(
|
||||||
"/{project_id}",
|
"/{project_id}",
|
||||||
response_model=ProjectResponse,
|
response_model=ProjectResponse,
|
||||||
|
|||||||
@@ -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
|
||||||
@@ -10,22 +10,13 @@ from pydantic import BaseModel, ConfigDict
|
|||||||
from sqlalchemy import select
|
from sqlalchemy import select
|
||||||
from sqlalchemy.ext.asyncio import AsyncSession
|
from sqlalchemy.ext.asyncio import AsyncSession
|
||||||
|
|
||||||
from src.auth.dependencies import get_current_user_id, get_db_session
|
from src.auth.dependencies import _get_user, get_current_user_id, get_db_session
|
||||||
from src.config import Settings
|
from src.config import Settings
|
||||||
from src.models.ssh_key import SSHKey
|
from src.models.ssh_key import SSHKey
|
||||||
from src.models.user import User
|
|
||||||
|
|
||||||
router = APIRouter(prefix="/ssh-keys", tags=["ssh-keys"])
|
router = APIRouter(prefix="/ssh-keys", tags=["ssh-keys"])
|
||||||
|
|
||||||
|
|
||||||
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
|
|
||||||
|
|
||||||
|
|
||||||
def _get_fernet() -> Fernet:
|
def _get_fernet() -> Fernet:
|
||||||
"""Generate a valid Fernet key from the session secret."""
|
"""Generate a valid Fernet key from the session secret."""
|
||||||
import base64
|
import base64
|
||||||
|
|||||||
+542
-77
@@ -1,15 +1,21 @@
|
|||||||
"""WebSocket terminal endpoint for tool instances."""
|
"""WebSocket terminal endpoint for tool instances."""
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
|
import json
|
||||||
import logging
|
import logging
|
||||||
import uuid
|
import uuid
|
||||||
|
from contextlib import suppress
|
||||||
|
|
||||||
from fastapi import APIRouter, Depends, HTTPException, WebSocket, WebSocketDisconnect, status
|
from fastapi import APIRouter, Depends, HTTPException, WebSocket, status
|
||||||
|
from sqlalchemy import select
|
||||||
from sqlalchemy.ext.asyncio import AsyncSession
|
from sqlalchemy.ext.asyncio import AsyncSession
|
||||||
|
from starlette.websockets import WebSocketDisconnect
|
||||||
|
|
||||||
from src.auth.dependencies import get_db_session
|
from src.auth.dependencies import get_current_user_id, get_db_session
|
||||||
|
from src.models.terminal_session import TerminalSessionModel
|
||||||
from src.models.tool_instance import ToolInstance
|
from src.models.tool_instance import ToolInstance
|
||||||
from src.services.terminal_manager import terminal_manager
|
from src.models.tool_type import ToolType
|
||||||
|
from src.services.terminal_manager import MaxSessionsExceededError, terminal_manager
|
||||||
|
|
||||||
router = APIRouter()
|
router = APIRouter()
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
@@ -18,33 +24,60 @@ logger = logging.getLogger(__name__)
|
|||||||
class SessionRef:
|
class SessionRef:
|
||||||
"""Mutable reference to a terminal session, allowing updates during reset."""
|
"""Mutable reference to a terminal session, allowing updates during reset."""
|
||||||
|
|
||||||
def __init__(self, session):
|
def __init__(self, session, slot_session_id: str | None = None):
|
||||||
self.session = session
|
self.session = session
|
||||||
|
self.slot_session_id = slot_session_id or session.session_id
|
||||||
|
|
||||||
|
|
||||||
@router.websocket(
|
@router.websocket(
|
||||||
"/ws/tool-instances/{instance_id}/terminal",
|
"/ws/tool-instances/{instance_id}/terminal",
|
||||||
)
|
)
|
||||||
async def terminal_websocket(
|
async def terminal_websocket_default(
|
||||||
websocket: WebSocket,
|
websocket: WebSocket,
|
||||||
instance_id: str,
|
instance_id: str,
|
||||||
db_session: AsyncSession = Depends(get_db_session),
|
db_session: AsyncSession = Depends(get_db_session),
|
||||||
) -> None:
|
) -> None:
|
||||||
"""WebSocket endpoint for terminal access to a tool instance.
|
"""WebSocket endpoint for terminal access (default session alias).
|
||||||
|
|
||||||
Provides an interactive terminal session inside a running tool instance container.
|
Backward-compatible route that maps to the default session.
|
||||||
Sessions persist across WebSocket disconnections.
|
"""
|
||||||
|
await _handle_terminal_websocket(websocket, instance_id, None, db_session)
|
||||||
|
|
||||||
|
|
||||||
|
@router.websocket(
|
||||||
|
"/ws/tool-instances/{instance_id}/terminal/{session_id}",
|
||||||
|
)
|
||||||
|
async def terminal_websocket_specific(
|
||||||
|
websocket: WebSocket,
|
||||||
|
instance_id: str,
|
||||||
|
session_id: str,
|
||||||
|
db_session: AsyncSession = Depends(get_db_session),
|
||||||
|
) -> None:
|
||||||
|
"""WebSocket endpoint for a specific terminal session."""
|
||||||
|
await _handle_terminal_websocket(websocket, instance_id, session_id, db_session)
|
||||||
|
|
||||||
|
|
||||||
|
async def _handle_terminal_websocket(
|
||||||
|
websocket: WebSocket,
|
||||||
|
instance_id: str,
|
||||||
|
target_session_id: str | None,
|
||||||
|
db_session: AsyncSession,
|
||||||
|
) -> None:
|
||||||
|
"""Shared WebSocket handler for terminal sessions.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
websocket: The WebSocket connection.
|
websocket: The WebSocket connection.
|
||||||
instance_id: UUID string of the tool instance.
|
instance_id: UUID string of the tool instance.
|
||||||
|
target_session_id: Specific session ID (slot key). None means default session.
|
||||||
db_session: Database session.
|
db_session: Database session.
|
||||||
|
|
||||||
Returns:
|
|
||||||
None. Communicates via WebSocket messages.
|
|
||||||
"""
|
"""
|
||||||
logger.info("Terminal WebSocket connection attempt for instance %s", instance_id)
|
logger.debug(
|
||||||
|
"Terminal WebSocket connection attempt for instance %s (session=%s)",
|
||||||
|
instance_id,
|
||||||
|
target_session_id or "default",
|
||||||
|
)
|
||||||
await websocket.accept()
|
await websocket.accept()
|
||||||
|
logger.debug("Terminal WebSocket accepted for instance %s", instance_id)
|
||||||
|
|
||||||
try:
|
try:
|
||||||
# Parse instance_id
|
# Parse instance_id
|
||||||
@@ -57,7 +90,9 @@ async def terminal_websocket(
|
|||||||
# Authenticate user from session cookie
|
# Authenticate user from session cookie
|
||||||
user_id = await _get_user_from_websocket(websocket, db_session)
|
user_id = await _get_user_from_websocket(websocket, db_session)
|
||||||
if user_id is None:
|
if user_id is None:
|
||||||
logger.warning("Unauthorized terminal access attempt for instance %s", instance_id)
|
logger.warning(
|
||||||
|
"Unauthorized terminal access attempt for instance %s", instance_id
|
||||||
|
)
|
||||||
await websocket.close(code=4003, reason="Unauthorized")
|
await websocket.close(code=4003, reason="Unauthorized")
|
||||||
return
|
return
|
||||||
|
|
||||||
@@ -69,59 +104,167 @@ async def terminal_websocket(
|
|||||||
return
|
return
|
||||||
|
|
||||||
if instance.owner_id != user_id:
|
if instance.owner_id != user_id:
|
||||||
logger.warning("Forbidden terminal access for instance %s by user %s", instance_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")
|
await websocket.close(code=4003, reason="Forbidden")
|
||||||
return
|
return
|
||||||
|
|
||||||
if instance.status != "running" or not instance.container_id:
|
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)
|
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")
|
await websocket.close(code=4004, reason="Instance not running")
|
||||||
return
|
return
|
||||||
|
|
||||||
|
logger.debug("Terminal auth passed for instance %s, user %s", instance_id, user_id)
|
||||||
|
|
||||||
|
# Verify the container actually exists (may have been removed/recreated)
|
||||||
|
from src.services.docker import get_container_status
|
||||||
|
|
||||||
|
container_status = get_container_status(instance.container_id)
|
||||||
|
if container_status["status"] == "not_found":
|
||||||
|
logger.error(
|
||||||
|
"Container %s for instance %s not found (may have been removed)",
|
||||||
|
instance.container_id,
|
||||||
|
instance_id,
|
||||||
|
)
|
||||||
|
await websocket.close(
|
||||||
|
code=4004, reason="Container not found — restart the tool instance"
|
||||||
|
)
|
||||||
|
return
|
||||||
|
|
||||||
|
# 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,
|
||||||
|
)
|
||||||
|
|
||||||
|
session = None
|
||||||
|
|
||||||
# Get or create terminal session
|
# Get or create terminal session
|
||||||
try:
|
try:
|
||||||
session = await terminal_manager.get_or_create_session(
|
if target_session_id is None:
|
||||||
instance_uuid,
|
# Default session alias
|
||||||
instance.container_id,
|
session = await terminal_manager.get_or_create_session(
|
||||||
|
instance_uuid,
|
||||||
|
instance.container_id,
|
||||||
|
startup_command=startup_command,
|
||||||
|
)
|
||||||
|
slot_session_id = "default"
|
||||||
|
else:
|
||||||
|
# Specific session
|
||||||
|
session = terminal_manager.get_session(
|
||||||
|
instance_id,
|
||||||
|
target_session_id,
|
||||||
|
)
|
||||||
|
if session is None:
|
||||||
|
# Session not in memory — may have been lost on server restart.
|
||||||
|
# Try to restore from the DB row.
|
||||||
|
db_row = await db_session.get(
|
||||||
|
TerminalSessionModel, uuid.UUID(target_session_id)
|
||||||
|
)
|
||||||
|
if (
|
||||||
|
db_row is not None
|
||||||
|
and db_row.instance_id == instance_uuid
|
||||||
|
and db_row.status != "closed"
|
||||||
|
):
|
||||||
|
logger.info(
|
||||||
|
"Restoring terminal session %s for instance %s from DB",
|
||||||
|
target_session_id,
|
||||||
|
instance_id,
|
||||||
|
)
|
||||||
|
session = await terminal_manager.create_session(
|
||||||
|
instance_uuid,
|
||||||
|
instance.container_id,
|
||||||
|
startup_command=startup_command,
|
||||||
|
name=db_row.name,
|
||||||
|
session_id=target_session_id,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
logger.warning(
|
||||||
|
"Session %s not found for instance %s",
|
||||||
|
target_session_id,
|
||||||
|
instance_id,
|
||||||
|
)
|
||||||
|
await websocket.close(code=4004, reason="Session not found")
|
||||||
|
return
|
||||||
|
# Determine slot key for reset scoping
|
||||||
|
key = terminal_manager._find_key_by_internal_id(
|
||||||
|
instance_id, session.session_id
|
||||||
|
)
|
||||||
|
slot_session_id = key[1] if key else target_session_id
|
||||||
|
|
||||||
|
logger.debug(
|
||||||
|
"Terminal session ready for instance %s (session_id=%s, slot=%s)",
|
||||||
|
instance_id,
|
||||||
|
session.session_id,
|
||||||
|
slot_session_id,
|
||||||
)
|
)
|
||||||
logger.info("Terminal session ready for instance %s (session_id=%s)", instance_id, session.session_id)
|
|
||||||
|
|
||||||
# Attach WebSocket to session
|
# Attach WebSocket to session
|
||||||
await terminal_manager.attach_websocket(session, websocket)
|
await terminal_manager.attach_websocket(session, websocket)
|
||||||
logger.info("WebSocket attached to session for instance %s", instance_id)
|
logger.debug("WebSocket attached to session for instance %s", instance_id)
|
||||||
|
|
||||||
# Send connected status
|
# Send connected status
|
||||||
await websocket.send_json({"type": "status", "status": "connected"})
|
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
|
# Use mutable session reference so loops can survive reset
|
||||||
session_ref = SessionRef(session)
|
session_ref = SessionRef(session, slot_session_id)
|
||||||
|
|
||||||
# Start I/O loops and heartbeat
|
# Start I/O loops and heartbeat
|
||||||
read_task = asyncio.create_task(_read_loop(session_ref, websocket))
|
read_task = asyncio.create_task(_read_loop(session_ref, websocket))
|
||||||
write_task = asyncio.create_task(_write_loop(session_ref, websocket, instance_id))
|
write_task = asyncio.create_task(
|
||||||
|
_write_loop(session_ref, websocket, instance_id)
|
||||||
|
)
|
||||||
heartbeat_task = asyncio.create_task(_heartbeat_loop(websocket))
|
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)
|
# Wait for either task to complete (indicating disconnect or error)
|
||||||
done, pending = await asyncio.wait(
|
done, pending = await asyncio.wait(
|
||||||
[read_task, write_task, heartbeat_task],
|
[read_task, write_task, heartbeat_task],
|
||||||
return_when=asyncio.FIRST_COMPLETED,
|
return_when=asyncio.FIRST_COMPLETED,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
logger.debug(
|
||||||
|
"Terminal loop completed for instance %s, done=%s",
|
||||||
|
instance_id,
|
||||||
|
len(done),
|
||||||
|
)
|
||||||
|
|
||||||
# Cancel remaining tasks
|
# Cancel remaining tasks
|
||||||
for task in pending:
|
for task in pending:
|
||||||
task.cancel()
|
task.cancel()
|
||||||
|
|
||||||
|
except WebSocketDisconnect:
|
||||||
|
logger.debug("WebSocket disconnected for instance %s", instance_id)
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
logger.error("Terminal session error for instance %s: %s", instance_id, str(exc), exc_info=True)
|
logger.error(
|
||||||
await websocket.close(code=4000, reason=f"Error: {exc}")
|
"Terminal session error for instance %s: %s",
|
||||||
|
instance_id,
|
||||||
|
str(exc),
|
||||||
|
exc_info=True,
|
||||||
|
)
|
||||||
|
with suppress(Exception):
|
||||||
|
await websocket.close(code=4000, reason=f"Error: {exc}")
|
||||||
finally:
|
finally:
|
||||||
# Detach WebSocket, don't kill session
|
# Detach WebSocket, don't kill session
|
||||||
try:
|
with suppress(Exception):
|
||||||
if 'session' in locals():
|
if session is not None:
|
||||||
await terminal_manager.detach_websocket(session, websocket)
|
await terminal_manager.detach_websocket(session, websocket)
|
||||||
logger.info("WebSocket detached from session for instance %s", instance_id)
|
logger.debug(
|
||||||
except Exception:
|
"WebSocket detached from session for instance %s", instance_id
|
||||||
pass
|
)
|
||||||
|
|
||||||
|
|
||||||
async def _read_loop(session_ref: SessionRef, websocket) -> None:
|
async def _read_loop(session_ref: SessionRef, websocket) -> None:
|
||||||
@@ -136,6 +279,8 @@ async def _read_loop(session_ref: SessionRef, websocket) -> None:
|
|||||||
if data:
|
if data:
|
||||||
try:
|
try:
|
||||||
await websocket.send_bytes(data)
|
await websocket.send_bytes(data)
|
||||||
|
except WebSocketDisconnect:
|
||||||
|
break
|
||||||
except Exception:
|
except Exception:
|
||||||
break
|
break
|
||||||
else:
|
else:
|
||||||
@@ -160,37 +305,54 @@ async def _write_loop(session_ref: SessionRef, websocket, instance_id: str) -> N
|
|||||||
text = message["text"]
|
text = message["text"]
|
||||||
if text.startswith("{"):
|
if text.startswith("{"):
|
||||||
# Control message (JSON)
|
# Control message (JSON)
|
||||||
import json
|
|
||||||
try:
|
try:
|
||||||
ctrl = json.loads(text)
|
ctrl = json.loads(text)
|
||||||
msg_type = ctrl.get("type")
|
msg_type = ctrl.get("type")
|
||||||
|
|
||||||
if msg_type == "resize":
|
if msg_type == "resize":
|
||||||
cols = ctrl.get("cols", 80)
|
cols = ctrl.get("cols", 80)
|
||||||
rows = ctrl.get("rows", 24)
|
rows = ctrl.get("rows", 24)
|
||||||
logger.info(f"Received resize message for instance {instance_id}: {cols}x{rows}")
|
logger.debug(
|
||||||
|
"Received resize message for instance %s: %sx%s",
|
||||||
|
instance_id,
|
||||||
|
cols,
|
||||||
|
rows,
|
||||||
|
)
|
||||||
await session.resize(cols, rows)
|
await session.resize(cols, rows)
|
||||||
elif msg_type == "reset":
|
elif msg_type == "reset":
|
||||||
# Reset terminal session
|
# Reset terminal session (scoped to current slot)
|
||||||
logger.info("Resetting terminal session for instance %s", session.instance_id)
|
logger.debug(
|
||||||
await websocket.send_json({"type": "status", "status": "resetting"})
|
"Resetting terminal session for instance %s (slot=%s)",
|
||||||
|
session.instance_id,
|
||||||
# Reset the session
|
session_ref.slot_session_id,
|
||||||
|
)
|
||||||
|
await websocket.send_json(
|
||||||
|
{"type": "status", "status": "resetting"}
|
||||||
|
)
|
||||||
|
|
||||||
|
# Reset the session scoped to its slot
|
||||||
new_session = await terminal_manager.reset_session(
|
new_session = await terminal_manager.reset_session(
|
||||||
session.instance_id,
|
session.instance_id,
|
||||||
session.container_id,
|
session.container_id,
|
||||||
|
startup_command=session.startup_command,
|
||||||
|
session_id=session_ref.slot_session_id,
|
||||||
|
name=session.name,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Update the mutable session reference so read_loop uses the new session
|
# Update the mutable session reference
|
||||||
session_ref.session = new_session
|
session_ref.session = new_session
|
||||||
|
|
||||||
# Attach to new session
|
# Attach to new session
|
||||||
await terminal_manager.attach_websocket(new_session, websocket)
|
await terminal_manager.attach_websocket(
|
||||||
await websocket.send_json({"type": "status", "status": "connected"})
|
new_session, websocket
|
||||||
|
)
|
||||||
|
await websocket.send_json(
|
||||||
|
{"type": "status", "status": "connected"}
|
||||||
|
)
|
||||||
|
|
||||||
# Continue the loop with the new session
|
# Continue the loop with the new session
|
||||||
continue
|
continue
|
||||||
|
|
||||||
except json.JSONDecodeError:
|
except json.JSONDecodeError:
|
||||||
# Not a valid JSON control message, treat as regular input
|
# Not a valid JSON control message, treat as regular input
|
||||||
await session.write_input(text.encode("utf-8"))
|
await session.write_input(text.encode("utf-8"))
|
||||||
@@ -216,51 +378,349 @@ async def _heartbeat_loop(websocket: WebSocket) -> None:
|
|||||||
pass
|
pass
|
||||||
|
|
||||||
|
|
||||||
@router.post(
|
async def _get_terminal_instance(
|
||||||
"/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,
|
instance_id: uuid.UUID,
|
||||||
db_session: AsyncSession = Depends(get_db_session),
|
user_id: uuid.UUID,
|
||||||
) -> dict:
|
db_session: AsyncSession,
|
||||||
"""Reset the terminal session for an instance.
|
) -> ToolInstance:
|
||||||
|
"""Fetch instance and validate auth, ownership, and running status.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
project_id: UUID of the project.
|
|
||||||
repo_id: UUID of the repository.
|
|
||||||
instance_id: UUID of the tool instance.
|
instance_id: UUID of the tool instance.
|
||||||
|
user_id: ID of the authenticated user.
|
||||||
|
db_session: Database session.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
The validated ToolInstance.
|
||||||
|
|
||||||
|
Raises:
|
||||||
|
HTTPException: If instance not found, not owned, or not 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.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_id:
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=status.HTTP_400_BAD_REQUEST, detail="Instance is not running"
|
||||||
|
)
|
||||||
|
|
||||||
|
return instance
|
||||||
|
|
||||||
|
|
||||||
|
@router.get(
|
||||||
|
"/instances/{instance_id}/terminal/sessions",
|
||||||
|
summary="List terminal sessions",
|
||||||
|
description="List terminal sessions for a tool instance with live WebSocket state.",
|
||||||
|
)
|
||||||
|
async def list_terminal_sessions(
|
||||||
|
instance_id: uuid.UUID,
|
||||||
|
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||||
|
db_session: AsyncSession = Depends(get_db_session),
|
||||||
|
) -> dict:
|
||||||
|
"""List terminal sessions for an instance.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
instance_id: UUID of the tool instance.
|
||||||
|
user_id: ID of the authenticated user.
|
||||||
|
db_session: Database session.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Dictionary with sessions list.
|
||||||
|
"""
|
||||||
|
await _get_terminal_instance(instance_id, user_id, db_session)
|
||||||
|
|
||||||
|
# Query active DB rows for this instance
|
||||||
|
result = await db_session.execute(
|
||||||
|
select(TerminalSessionModel)
|
||||||
|
.where(TerminalSessionModel.instance_id == instance_id)
|
||||||
|
.where(TerminalSessionModel.status != "closed")
|
||||||
|
.order_by(TerminalSessionModel.created_at.asc())
|
||||||
|
)
|
||||||
|
db_rows = result.scalars().all()
|
||||||
|
|
||||||
|
# Build response with live has_websockets flag.
|
||||||
|
# Include DB rows even without in-memory counterparts (e.g. after
|
||||||
|
# server restart) so the frontend can display tabs and reconnect.
|
||||||
|
sessions = []
|
||||||
|
for row in db_rows:
|
||||||
|
live_session = terminal_manager.get_session(str(instance_id), str(row.id))
|
||||||
|
sessions.append(
|
||||||
|
{
|
||||||
|
"id": str(row.id),
|
||||||
|
"name": row.name,
|
||||||
|
"status": row.status,
|
||||||
|
"has_websockets": live_session.has_websockets()
|
||||||
|
if live_session
|
||||||
|
else False,
|
||||||
|
"created_at": row.created_at.isoformat() if row.created_at else None,
|
||||||
|
"last_activity_at": row.last_activity_at.isoformat()
|
||||||
|
if row.last_activity_at
|
||||||
|
else None,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
|
return {"sessions": sessions}
|
||||||
|
|
||||||
|
|
||||||
|
@router.post(
|
||||||
|
"/instances/{instance_id}/terminal/sessions",
|
||||||
|
summary="Create terminal session",
|
||||||
|
description="Create a new terminal session for a running tool instance.",
|
||||||
|
status_code=status.HTTP_201_CREATED,
|
||||||
|
)
|
||||||
|
async def create_terminal_session(
|
||||||
|
instance_id: uuid.UUID,
|
||||||
|
data: dict,
|
||||||
|
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||||
|
db_session: AsyncSession = Depends(get_db_session),
|
||||||
|
) -> dict:
|
||||||
|
"""Create a new terminal session.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
instance_id: UUID of the tool instance.
|
||||||
|
data: Request body with optional name.
|
||||||
|
user_id: ID of the authenticated user.
|
||||||
|
db_session: Database session.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Dictionary with new session details.
|
||||||
|
|
||||||
|
Raises:
|
||||||
|
HTTPException: 409 if max sessions reached.
|
||||||
|
"""
|
||||||
|
instance = await _get_terminal_instance(instance_id, user_id, db_session)
|
||||||
|
assert instance.container_id is not None
|
||||||
|
|
||||||
|
# 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
|
||||||
|
|
||||||
|
name = data.get("name")
|
||||||
|
|
||||||
|
try:
|
||||||
|
session = await terminal_manager.create_session(
|
||||||
|
instance_id,
|
||||||
|
instance.container_id,
|
||||||
|
startup_command=startup_command,
|
||||||
|
name=name,
|
||||||
|
)
|
||||||
|
except MaxSessionsExceededError:
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=status.HTTP_409_CONFLICT,
|
||||||
|
detail="Maximum of 5 terminal sessions reached for this instance",
|
||||||
|
) from None
|
||||||
|
|
||||||
|
return {
|
||||||
|
"id": session.session_id,
|
||||||
|
"name": session.name,
|
||||||
|
"status": session.status,
|
||||||
|
"created_at": session.last_activity,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@router.delete(
|
||||||
|
"/instances/{instance_id}/terminal/sessions/{session_id}",
|
||||||
|
summary="Close terminal session",
|
||||||
|
description="Close a specific terminal session.",
|
||||||
|
)
|
||||||
|
async def close_terminal_session(
|
||||||
|
instance_id: uuid.UUID,
|
||||||
|
session_id: str,
|
||||||
|
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||||
|
db_session: AsyncSession = Depends(get_db_session),
|
||||||
|
) -> dict:
|
||||||
|
"""Close a terminal session.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
instance_id: UUID of the tool instance.
|
||||||
|
session_id: ID of the session to close.
|
||||||
|
user_id: ID of the authenticated user.
|
||||||
|
db_session: Database session.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Dictionary with closure status.
|
||||||
|
"""
|
||||||
|
await _get_terminal_instance(instance_id, user_id, db_session)
|
||||||
|
|
||||||
|
# Find the session by internal ID to determine its slot key
|
||||||
|
key = terminal_manager._find_key_by_internal_id(str(instance_id), session_id)
|
||||||
|
if (
|
||||||
|
key is None
|
||||||
|
and terminal_manager.get_session(str(instance_id), session_id) is not None
|
||||||
|
):
|
||||||
|
key = (str(instance_id), session_id)
|
||||||
|
|
||||||
|
if key is None:
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=status.HTTP_404_NOT_FOUND, detail="Session not found"
|
||||||
|
)
|
||||||
|
|
||||||
|
await terminal_manager.close_session(key[0], key[1])
|
||||||
|
|
||||||
|
return {"status": "closed", "session_id": session_id}
|
||||||
|
|
||||||
|
|
||||||
|
@router.post(
|
||||||
|
"/instances/{instance_id}/terminal/sessions/{session_id}/reset",
|
||||||
|
summary="Reset terminal session",
|
||||||
|
description="Reset a specific terminal session, killing the current shell and starting fresh.",
|
||||||
|
)
|
||||||
|
async def reset_specific_terminal_session(
|
||||||
|
instance_id: uuid.UUID,
|
||||||
|
session_id: str,
|
||||||
|
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||||
|
db_session: AsyncSession = Depends(get_db_session),
|
||||||
|
) -> dict:
|
||||||
|
"""Reset a specific terminal session.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
instance_id: UUID of the tool instance.
|
||||||
|
session_id: ID of the session to reset.
|
||||||
|
user_id: ID of the authenticated user.
|
||||||
|
db_session: Database session.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Dictionary with reset session details.
|
||||||
|
"""
|
||||||
|
instance = await _get_terminal_instance(instance_id, user_id, db_session)
|
||||||
|
assert instance.container_id is not None
|
||||||
|
|
||||||
|
# Determine slot key for reset
|
||||||
|
key = terminal_manager._find_key_by_internal_id(str(instance_id), session_id)
|
||||||
|
if (
|
||||||
|
key is None
|
||||||
|
and terminal_manager.get_session(str(instance_id), session_id) is not None
|
||||||
|
):
|
||||||
|
key = (str(instance_id), session_id)
|
||||||
|
|
||||||
|
if key is None:
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=status.HTTP_404_NOT_FOUND, detail="Session not found"
|
||||||
|
)
|
||||||
|
|
||||||
|
# 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
|
||||||
|
|
||||||
|
# Preserve name if possible
|
||||||
|
live_session = terminal_manager.get_session(str(instance_id), session_id)
|
||||||
|
name = live_session.name if live_session else None
|
||||||
|
|
||||||
|
new_session = await terminal_manager.reset_session(
|
||||||
|
instance_id,
|
||||||
|
instance.container_id,
|
||||||
|
startup_command=startup_command,
|
||||||
|
session_id=key[1],
|
||||||
|
name=name,
|
||||||
|
)
|
||||||
|
|
||||||
|
return {
|
||||||
|
"id": new_session.session_id,
|
||||||
|
"name": new_session.name,
|
||||||
|
"status": new_session.status,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@router.post(
|
||||||
|
"/instances/{instance_id}/terminal/sessions/{session_id}/rename",
|
||||||
|
summary="Rename terminal session",
|
||||||
|
description="Rename a specific terminal session.",
|
||||||
|
)
|
||||||
|
async def rename_terminal_session(
|
||||||
|
instance_id: uuid.UUID,
|
||||||
|
session_id: str,
|
||||||
|
data: dict,
|
||||||
|
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||||
|
db_session: AsyncSession = Depends(get_db_session),
|
||||||
|
) -> dict:
|
||||||
|
"""Rename a terminal session.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
instance_id: UUID of the tool instance.
|
||||||
|
session_id: ID of the session to rename.
|
||||||
|
data: Request body with new name.
|
||||||
|
user_id: ID of the authenticated user.
|
||||||
|
db_session: Database session.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Dictionary with updated session details.
|
||||||
|
"""
|
||||||
|
await _get_terminal_instance(instance_id, user_id, db_session)
|
||||||
|
|
||||||
|
new_name = data.get("name")
|
||||||
|
if not new_name or not isinstance(new_name, str):
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=status.HTTP_400_BAD_REQUEST, detail="Name is required"
|
||||||
|
)
|
||||||
|
|
||||||
|
# Update in-memory session name if live
|
||||||
|
live_session = terminal_manager.get_session(str(instance_id), session_id)
|
||||||
|
if live_session:
|
||||||
|
live_session.name = new_name
|
||||||
|
|
||||||
|
# Update DB row
|
||||||
|
db_row = await db_session.get(TerminalSessionModel, uuid.UUID(session_id))
|
||||||
|
if db_row is None:
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=status.HTTP_404_NOT_FOUND, detail="Session not found"
|
||||||
|
)
|
||||||
|
|
||||||
|
db_row.name = new_name
|
||||||
|
await db_session.commit()
|
||||||
|
|
||||||
|
return {"id": str(db_row.id), "name": new_name}
|
||||||
|
|
||||||
|
|
||||||
|
@router.post(
|
||||||
|
"/instances/{instance_id}/terminal/reset",
|
||||||
|
summary="Reset terminal session (legacy alias)",
|
||||||
|
description="Reset the default terminal session for a tool instance. Preserved for backward compatibility.",
|
||||||
|
)
|
||||||
|
async def reset_terminal_session(
|
||||||
|
instance_id: uuid.UUID,
|
||||||
|
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||||
|
db_session: AsyncSession = Depends(get_db_session),
|
||||||
|
) -> dict:
|
||||||
|
"""Reset the default terminal session for an instance (legacy alias).
|
||||||
|
|
||||||
|
Args:
|
||||||
|
instance_id: UUID of the tool instance.
|
||||||
|
user_id: ID of the authenticated user.
|
||||||
db_session: Database session.
|
db_session: Database session.
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Dictionary with status message.
|
Dictionary with status message.
|
||||||
"""
|
"""
|
||||||
# Get instance and verify it exists and is running
|
instance = await _get_terminal_instance(instance_id, user_id, db_session)
|
||||||
instance = await db_session.get(ToolInstance, instance_id)
|
assert instance.container_id is not None
|
||||||
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:
|
# Fetch tool type to get startup_command
|
||||||
raise HTTPException(
|
tool_type = await db_session.get(ToolType, instance.tool_type_id)
|
||||||
status_code=status.HTTP_400_BAD_REQUEST,
|
startup_command = tool_type.startup_command if tool_type else None
|
||||||
detail="Instance is not running"
|
|
||||||
)
|
|
||||||
|
|
||||||
try:
|
try:
|
||||||
# Reset the session
|
# Reset the default session
|
||||||
new_session = await terminal_manager.reset_session(
|
new_session = await terminal_manager.reset_session(
|
||||||
instance_id,
|
instance_id,
|
||||||
instance.container_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)
|
logger.info(
|
||||||
|
"Terminal session reset for instance %s (new session_id=%s)",
|
||||||
|
instance_id,
|
||||||
|
new_session.session_id,
|
||||||
|
)
|
||||||
|
|
||||||
return {
|
return {
|
||||||
"status": "success",
|
"status": "success",
|
||||||
"message": "Terminal session reset successfully",
|
"message": "Terminal session reset successfully",
|
||||||
@@ -268,11 +728,16 @@ async def reset_terminal_session(
|
|||||||
"session_id": new_session.session_id,
|
"session_id": new_session.session_id,
|
||||||
}
|
}
|
||||||
except Exception as exc:
|
except Exception as exc:
|
||||||
logger.error("Failed to reset terminal session for instance %s: %s", instance_id, str(exc), exc_info=True)
|
logger.error(
|
||||||
|
"Failed to reset terminal session for instance %s: %s",
|
||||||
|
instance_id,
|
||||||
|
str(exc),
|
||||||
|
exc_info=True,
|
||||||
|
)
|
||||||
raise HTTPException(
|
raise HTTPException(
|
||||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||||
detail=f"Failed to reset terminal session: {exc}"
|
detail=f"Failed to reset terminal session: {exc}",
|
||||||
)
|
) from exc
|
||||||
|
|
||||||
|
|
||||||
async def _get_user_from_websocket(
|
async def _get_user_from_websocket(
|
||||||
|
|||||||
@@ -1,322 +0,0 @@
|
|||||||
"""Tool configuration 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.auth.dependencies import get_current_user_id, get_db_session
|
|
||||||
from src.models.tool_config import ToolConfig
|
|
||||||
from src.models.tool_type import ToolType
|
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
|
||||||
|
|
||||||
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:
|
|
||||||
if v is None:
|
|
||||||
return v
|
|
||||||
if not isinstance(v, dict):
|
|
||||||
raise ValueError("environment_variables must be a JSON object")
|
|
||||||
return v
|
|
||||||
|
|
||||||
@field_validator("volumes")
|
|
||||||
@classmethod
|
|
||||||
def validate_volumes(cls, v: list | None) -> list | None:
|
|
||||||
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
|
|
||||||
|
|
||||||
|
|
||||||
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:
|
|
||||||
if v is None:
|
|
||||||
return v
|
|
||||||
if not isinstance(v, dict):
|
|
||||||
raise ValueError("environment_variables must be a JSON object")
|
|
||||||
return v
|
|
||||||
|
|
||||||
@field_validator("volumes")
|
|
||||||
@classmethod
|
|
||||||
def validate_volumes(cls, v: list | None) -> list | None:
|
|
||||||
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
|
|
||||||
|
|
||||||
|
|
||||||
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()
|
|
||||||
@@ -0,0 +1,424 @@
|
|||||||
|
"""Tool definition API endpoints."""
|
||||||
|
|
||||||
|
import logging
|
||||||
|
import uuid
|
||||||
|
|
||||||
|
from fastapi import APIRouter, Depends, HTTPException, status
|
||||||
|
from pydantic import BaseModel, Field
|
||||||
|
from sqlalchemy import select
|
||||||
|
from sqlalchemy.ext.asyncio import AsyncSession
|
||||||
|
|
||||||
|
from src.auth.dependencies import get_current_user_id, get_db_session
|
||||||
|
from src.models.tool_definition_manifest import ToolDefinitionManifest
|
||||||
|
from src.models.tool_type import ToolType
|
||||||
|
from src.services.manifest_compiler import (
|
||||||
|
compile_compose,
|
||||||
|
compile_dockerfile,
|
||||||
|
compile_entrypoint,
|
||||||
|
compute_image_tag,
|
||||||
|
deep_merge,
|
||||||
|
resolve_base,
|
||||||
|
)
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
router = APIRouter(prefix="/tool-definitions", tags=["tool-definitions"])
|
||||||
|
|
||||||
|
|
||||||
|
class CreateToolDefinitionRequest(BaseModel):
|
||||||
|
"""Request body for creating a tool definition manifest."""
|
||||||
|
|
||||||
|
model_config = {"extra": "ignore"}
|
||||||
|
|
||||||
|
name: str = Field(description="Unique identifier (kebab-case)")
|
||||||
|
display_name: str = Field(description="Human-readable name")
|
||||||
|
description: str | None = Field(default=None)
|
||||||
|
category: str = Field(default="development")
|
||||||
|
interface_type: str = Field(default="terminal", description="web or terminal")
|
||||||
|
base_image: str | None = Field(default=None, description="Direct base image")
|
||||||
|
base_definition_id: str | None = Field(
|
||||||
|
default=None, description="Reference to a base definition"
|
||||||
|
)
|
||||||
|
base_version: str = Field(default="latest")
|
||||||
|
manifest: dict = Field(description="The full manifest JSON")
|
||||||
|
|
||||||
|
|
||||||
|
class UpdateToolDefinitionRequest(BaseModel):
|
||||||
|
"""Request body for updating a tool definition manifest."""
|
||||||
|
|
||||||
|
model_config = {"extra": "ignore"}
|
||||||
|
|
||||||
|
display_name: str | None = Field(default=None)
|
||||||
|
description: str | None = Field(default=None)
|
||||||
|
category: str | None = Field(default=None)
|
||||||
|
manifest: dict | None = Field(default=None)
|
||||||
|
base_version: str | None = Field(default=None)
|
||||||
|
|
||||||
|
|
||||||
|
@router.post(
|
||||||
|
"",
|
||||||
|
summary="Create tool definition",
|
||||||
|
description="Create a new tool definition manifest.",
|
||||||
|
)
|
||||||
|
async def create_tool_definition(
|
||||||
|
data: CreateToolDefinitionRequest,
|
||||||
|
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||||
|
session: AsyncSession = Depends(get_db_session),
|
||||||
|
) -> dict:
|
||||||
|
"""Create a new tool definition manifest.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
data: Manifest data.
|
||||||
|
user_id: Authenticated user ID.
|
||||||
|
session: Database session.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Dictionary with created definition details.
|
||||||
|
"""
|
||||||
|
# Validate base reference
|
||||||
|
if not data.base_image and not data.base_definition_id:
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=status.HTTP_400_BAD_REQUEST,
|
||||||
|
detail="Either base_image or base_definition_id is required",
|
||||||
|
)
|
||||||
|
|
||||||
|
base_def_id = None
|
||||||
|
if data.base_definition_id:
|
||||||
|
try:
|
||||||
|
base_def_id = uuid.UUID(data.base_definition_id)
|
||||||
|
except ValueError:
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=status.HTTP_400_BAD_REQUEST,
|
||||||
|
detail=f"Invalid base_definition_id: {data.base_definition_id}",
|
||||||
|
)
|
||||||
|
|
||||||
|
base_def = await session.get(ToolDefinitionManifest, base_def_id)
|
||||||
|
if not base_def:
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=status.HTTP_404_NOT_FOUND,
|
||||||
|
detail=f"Base definition not found: {data.base_definition_id}",
|
||||||
|
)
|
||||||
|
if not base_def.is_base:
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=status.HTTP_400_BAD_REQUEST,
|
||||||
|
detail="Referenced definition is not a base definition",
|
||||||
|
)
|
||||||
|
|
||||||
|
# Check name uniqueness
|
||||||
|
existing = await session.execute(
|
||||||
|
select(ToolDefinitionManifest).where(ToolDefinitionManifest.name == data.name)
|
||||||
|
)
|
||||||
|
if existing.scalar_one_or_none():
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=status.HTTP_409_CONFLICT,
|
||||||
|
detail=f"Tool definition '{data.name}' already exists",
|
||||||
|
)
|
||||||
|
|
||||||
|
definition = ToolDefinitionManifest(
|
||||||
|
name=data.name,
|
||||||
|
display_name=data.display_name,
|
||||||
|
description=data.description,
|
||||||
|
category=data.category,
|
||||||
|
interface_type=data.interface_type,
|
||||||
|
base_image=data.base_image,
|
||||||
|
base_definition_id=base_def_id,
|
||||||
|
base_version=data.base_version,
|
||||||
|
manifest=data.manifest,
|
||||||
|
created_by_id=user_id,
|
||||||
|
)
|
||||||
|
|
||||||
|
session.add(definition)
|
||||||
|
await session.commit()
|
||||||
|
await session.refresh(definition)
|
||||||
|
|
||||||
|
logger.info("Created tool definition %s (%s)", definition.id, definition.name)
|
||||||
|
|
||||||
|
return {
|
||||||
|
"id": str(definition.id),
|
||||||
|
"name": definition.name,
|
||||||
|
"display_name": definition.display_name,
|
||||||
|
"description": definition.description,
|
||||||
|
"category": definition.category,
|
||||||
|
"interface_type": definition.interface_type,
|
||||||
|
"base_image": definition.base_image,
|
||||||
|
"base_definition_id": str(definition.base_definition_id)
|
||||||
|
if definition.base_definition_id
|
||||||
|
else None,
|
||||||
|
"base_version": definition.base_version,
|
||||||
|
"manifest": definition.manifest,
|
||||||
|
"is_base": definition.is_base,
|
||||||
|
"created_at": definition.created_at.isoformat(),
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@router.get(
|
||||||
|
"",
|
||||||
|
summary="List tool definitions",
|
||||||
|
description="List all tool definition manifests.",
|
||||||
|
)
|
||||||
|
async def list_tool_definitions(
|
||||||
|
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||||
|
session: AsyncSession = Depends(get_db_session),
|
||||||
|
include_bases: bool = True,
|
||||||
|
) -> dict:
|
||||||
|
"""List all tool definition manifests.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
user_id: Authenticated user ID.
|
||||||
|
session: Database session.
|
||||||
|
include_bases: Whether to include base definitions.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Dictionary containing list of definitions.
|
||||||
|
"""
|
||||||
|
query = select(ToolDefinitionManifest)
|
||||||
|
if not include_bases:
|
||||||
|
query = query.where(ToolDefinitionManifest.is_base == False)
|
||||||
|
|
||||||
|
result = await session.execute(
|
||||||
|
query.order_by(ToolDefinitionManifest.created_at.desc())
|
||||||
|
)
|
||||||
|
definitions = result.scalars().all()
|
||||||
|
|
||||||
|
return {
|
||||||
|
"definitions": [
|
||||||
|
{
|
||||||
|
"id": str(d.id),
|
||||||
|
"name": d.name,
|
||||||
|
"display_name": d.display_name,
|
||||||
|
"description": d.description,
|
||||||
|
"category": d.category,
|
||||||
|
"interface_type": d.interface_type,
|
||||||
|
"is_base": d.is_base,
|
||||||
|
"base_image": d.base_image,
|
||||||
|
"base_definition_id": str(d.base_definition_id)
|
||||||
|
if d.base_definition_id
|
||||||
|
else None,
|
||||||
|
"base_version": d.base_version,
|
||||||
|
"version": d.version,
|
||||||
|
"created_at": d.created_at.isoformat(),
|
||||||
|
}
|
||||||
|
for d in definitions
|
||||||
|
]
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@router.get(
|
||||||
|
"/{definition_id}",
|
||||||
|
summary="Get tool definition",
|
||||||
|
description="Get a specific tool definition manifest.",
|
||||||
|
)
|
||||||
|
async def get_tool_definition(
|
||||||
|
definition_id: uuid.UUID,
|
||||||
|
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||||
|
session: AsyncSession = Depends(get_db_session),
|
||||||
|
) -> dict:
|
||||||
|
"""Get a specific tool definition manifest.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
definition_id: UUID of the definition.
|
||||||
|
user_id: Authenticated user ID.
|
||||||
|
session: Database session.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Dictionary with definition details.
|
||||||
|
"""
|
||||||
|
definition = await session.get(ToolDefinitionManifest, definition_id)
|
||||||
|
if not definition:
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=status.HTTP_404_NOT_FOUND,
|
||||||
|
detail=f"Tool definition not found: {definition_id}",
|
||||||
|
)
|
||||||
|
|
||||||
|
return {
|
||||||
|
"id": str(definition.id),
|
||||||
|
"name": definition.name,
|
||||||
|
"display_name": definition.display_name,
|
||||||
|
"description": definition.description,
|
||||||
|
"category": definition.category,
|
||||||
|
"interface_type": definition.interface_type,
|
||||||
|
"base_image": definition.base_image,
|
||||||
|
"base_definition_id": str(definition.base_definition_id)
|
||||||
|
if definition.base_definition_id
|
||||||
|
else None,
|
||||||
|
"base_version": definition.base_version,
|
||||||
|
"manifest": definition.manifest,
|
||||||
|
"dockerfile_cache": definition.dockerfile_cache,
|
||||||
|
"compose_cache": definition.compose_cache,
|
||||||
|
"version": definition.version,
|
||||||
|
"is_base": definition.is_base,
|
||||||
|
"created_at": definition.created_at.isoformat(),
|
||||||
|
"updated_at": definition.updated_at.isoformat(),
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@router.put(
|
||||||
|
"/{definition_id}",
|
||||||
|
summary="Update tool definition",
|
||||||
|
description="Update a tool definition manifest.",
|
||||||
|
)
|
||||||
|
async def update_tool_definition(
|
||||||
|
definition_id: uuid.UUID,
|
||||||
|
data: UpdateToolDefinitionRequest,
|
||||||
|
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||||
|
session: AsyncSession = Depends(get_db_session),
|
||||||
|
) -> dict:
|
||||||
|
"""Update a tool definition manifest.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
definition_id: UUID of the definition.
|
||||||
|
data: Update data.
|
||||||
|
user_id: Authenticated user ID.
|
||||||
|
session: Database session.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Dictionary with updated definition details.
|
||||||
|
"""
|
||||||
|
definition = await session.get(ToolDefinitionManifest, definition_id)
|
||||||
|
if not definition:
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=status.HTTP_404_NOT_FOUND,
|
||||||
|
detail=f"Tool definition not found: {definition_id}",
|
||||||
|
)
|
||||||
|
|
||||||
|
if data.display_name is not None:
|
||||||
|
definition.display_name = data.display_name
|
||||||
|
if data.description is not None:
|
||||||
|
definition.description = data.description
|
||||||
|
if data.category is not None:
|
||||||
|
definition.category = data.category
|
||||||
|
if data.manifest is not None:
|
||||||
|
definition.manifest = data.manifest
|
||||||
|
if data.base_version is not None:
|
||||||
|
definition.base_version = data.base_version
|
||||||
|
|
||||||
|
await session.commit()
|
||||||
|
await session.refresh(definition)
|
||||||
|
|
||||||
|
logger.info("Updated tool definition %s (%s)", definition.id, definition.name)
|
||||||
|
|
||||||
|
return {
|
||||||
|
"id": str(definition.id),
|
||||||
|
"name": definition.name,
|
||||||
|
"display_name": definition.display_name,
|
||||||
|
"manifest": definition.manifest,
|
||||||
|
"updated_at": definition.updated_at.isoformat(),
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@router.delete(
|
||||||
|
"/{definition_id}",
|
||||||
|
summary="Delete tool definition",
|
||||||
|
description="Delete a tool definition manifest.",
|
||||||
|
)
|
||||||
|
async def delete_tool_definition(
|
||||||
|
definition_id: uuid.UUID,
|
||||||
|
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||||
|
session: AsyncSession = Depends(get_db_session),
|
||||||
|
) -> dict:
|
||||||
|
"""Delete a tool definition manifest.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
definition_id: UUID of the definition.
|
||||||
|
user_id: Authenticated user ID.
|
||||||
|
session: Database session.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Dictionary with deletion status.
|
||||||
|
"""
|
||||||
|
definition = await session.get(ToolDefinitionManifest, definition_id)
|
||||||
|
if not definition:
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=status.HTTP_404_NOT_FOUND,
|
||||||
|
detail=f"Tool definition not found: {definition_id}",
|
||||||
|
)
|
||||||
|
|
||||||
|
# Check if any tool types reference this manifest
|
||||||
|
result = await session.execute(
|
||||||
|
select(ToolType).where(ToolType.manifest_id == definition_id)
|
||||||
|
)
|
||||||
|
referencing = result.scalars().all()
|
||||||
|
if referencing:
|
||||||
|
tool_names = ", ".join(t.name for t in referencing)
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=status.HTTP_409_CONFLICT,
|
||||||
|
detail=f"Cannot delete: referenced by tool types: {tool_names}",
|
||||||
|
)
|
||||||
|
|
||||||
|
await session.delete(definition)
|
||||||
|
await session.commit()
|
||||||
|
|
||||||
|
logger.info("Deleted tool definition %s (%s)", definition.id, definition.name)
|
||||||
|
|
||||||
|
return {"status": "deleted", "id": str(definition_id)}
|
||||||
|
|
||||||
|
|
||||||
|
@router.post(
|
||||||
|
"/{definition_id}/compile",
|
||||||
|
summary="Compile tool definition",
|
||||||
|
description="Compile a manifest to Dockerfile + Compose preview without building.",
|
||||||
|
)
|
||||||
|
async def compile_tool_definition(
|
||||||
|
definition_id: uuid.UUID,
|
||||||
|
user_id: uuid.UUID = Depends(get_current_user_id),
|
||||||
|
session: AsyncSession = Depends(get_db_session),
|
||||||
|
) -> dict:
|
||||||
|
"""Compile a manifest to Dockerfile + Compose preview.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
definition_id: UUID of the definition.
|
||||||
|
user_id: Authenticated user ID.
|
||||||
|
session: Database session.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Dictionary with compiled Dockerfile, Compose, and image tag.
|
||||||
|
"""
|
||||||
|
definition = await session.get(ToolDefinitionManifest, definition_id)
|
||||||
|
if not definition:
|
||||||
|
raise HTTPException(
|
||||||
|
status_code=status.HTTP_404_NOT_FOUND,
|
||||||
|
detail=f"Tool definition not found: {definition_id}",
|
||||||
|
)
|
||||||
|
|
||||||
|
manifest = dict(definition.manifest)
|
||||||
|
|
||||||
|
# Resolve base if referenced
|
||||||
|
if definition.base_definition_id:
|
||||||
|
base_def = await session.get(
|
||||||
|
ToolDefinitionManifest, definition.base_definition_id
|
||||||
|
)
|
||||||
|
if base_def:
|
||||||
|
base_manifest = dict(base_def.manifest)
|
||||||
|
manifest = resolve_base(deep_merge(base_manifest, manifest))
|
||||||
|
|
||||||
|
# Compile
|
||||||
|
dockerfile = compile_dockerfile(manifest)
|
||||||
|
entrypoint = compile_entrypoint(manifest)
|
||||||
|
image_tag = compute_image_tag(definition.name, manifest)
|
||||||
|
|
||||||
|
# Dummy compose with placeholder variables
|
||||||
|
dummy_vars = {
|
||||||
|
"IMAGE_TAG": image_tag,
|
||||||
|
"INSTANCE_NAME": f"{definition.name}-preview",
|
||||||
|
"INSTANCE_DIR": "/data/instances/preview",
|
||||||
|
"REPO_PATH": "/data/repos/preview",
|
||||||
|
"SSH_PATH": "/data/instances/preview/.ssh",
|
||||||
|
"TOOL_PORT": "8080",
|
||||||
|
"EXTRA_ENV": {},
|
||||||
|
"EXTRA_VOLUMES": [],
|
||||||
|
}
|
||||||
|
compose = compile_compose(manifest, dummy_vars)
|
||||||
|
|
||||||
|
# Update cache
|
||||||
|
definition.dockerfile_cache = dockerfile
|
||||||
|
definition.compose_cache = compose
|
||||||
|
await session.commit()
|
||||||
|
|
||||||
|
return {
|
||||||
|
"id": str(definition.id),
|
||||||
|
"name": definition.name,
|
||||||
|
"dockerfile": dockerfile,
|
||||||
|
"entrypoint": entrypoint,
|
||||||
|
"compose": compose,
|
||||||
|
"image_tag": image_tag,
|
||||||
|
}
|
||||||
+1433
-259
File diff suppressed because it is too large
Load Diff
+148
-185
@@ -1,33 +1,23 @@
|
|||||||
import re
|
|
||||||
import uuid
|
import uuid
|
||||||
from datetime import datetime
|
from datetime import datetime
|
||||||
|
|
||||||
import yaml
|
|
||||||
from fastapi import APIRouter, Depends, HTTPException, status
|
from fastapi import APIRouter, Depends, HTTPException, status
|
||||||
from pydantic import BaseModel, ConfigDict, field_validator, model_validator
|
from pydantic import BaseModel, ConfigDict, field_validator, model_validator
|
||||||
from sqlalchemy import select
|
from sqlalchemy import select
|
||||||
from sqlalchemy.ext.asyncio import AsyncSession
|
from sqlalchemy.ext.asyncio import AsyncSession
|
||||||
|
|
||||||
|
from src.api.tool_types_validation import (
|
||||||
def _sanitize_template_vars(template: str) -> str:
|
check_port_exposed,
|
||||||
"""Replace template variables like {{VAR}} with placeholders to avoid YAML parsing errors."""
|
validate_compose_yaml,
|
||||||
return re.sub(r"\{\{[A-Za-z_][A-Za-z0-9_]*\}\}", "__PLACEHOLDER__", template)
|
validate_required_variables,
|
||||||
|
)
|
||||||
from src.auth.dependencies import get_current_user_id, get_db_session
|
from src.auth.dependencies import _get_user, get_current_user_id, get_db_session
|
||||||
from src.models.tool_type import ToolType
|
from src.models.tool_type import ToolType
|
||||||
from src.models.user import User
|
from src.models.user import User
|
||||||
|
|
||||||
router = APIRouter(prefix="/tool-types", tags=["tool-types"])
|
router = APIRouter(prefix="/tool-types", tags=["tool-types"])
|
||||||
|
|
||||||
|
|
||||||
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 _require_admin(user: User) -> None:
|
async def _require_admin(user: User) -> None:
|
||||||
"""Check if user has admin privileges.
|
"""Check if user has admin privileges.
|
||||||
|
|
||||||
@@ -45,10 +35,12 @@ class ToolTypeCreate(BaseModel):
|
|||||||
description: str | None = None
|
description: str | None = None
|
||||||
default_port: int = 0
|
default_port: int = 0
|
||||||
definition_type: str = "compose"
|
definition_type: str = "compose"
|
||||||
|
manifest_id: uuid.UUID | None = None
|
||||||
compose_template: str | None = None
|
compose_template: str | None = None
|
||||||
dockerfile_template: str | None = None
|
dockerfile_template: str | None = None
|
||||||
build_context: dict | None = None
|
build_context: dict | None = None
|
||||||
readiness_probe: dict | None = None
|
readiness_probe: dict | None = None
|
||||||
|
startup_command: str | None = None
|
||||||
required_variables: list[str] = []
|
required_variables: list[str] = []
|
||||||
category: str = "other"
|
category: str = "other"
|
||||||
interface_type: str = "web"
|
interface_type: str = "web"
|
||||||
@@ -57,8 +49,10 @@ class ToolTypeCreate(BaseModel):
|
|||||||
@field_validator("definition_type")
|
@field_validator("definition_type")
|
||||||
@classmethod
|
@classmethod
|
||||||
def validate_definition_type(cls, v: str) -> str:
|
def validate_definition_type(cls, v: str) -> str:
|
||||||
if v not in ("compose", "dockerfile"):
|
if v not in ("compose", "dockerfile", "manifest"):
|
||||||
raise ValueError("definition_type must be 'compose' or 'dockerfile'")
|
raise ValueError(
|
||||||
|
"definition_type must be 'compose', 'dockerfile', or 'manifest'"
|
||||||
|
)
|
||||||
return v
|
return v
|
||||||
|
|
||||||
@field_validator("compose_template")
|
@field_validator("compose_template")
|
||||||
@@ -67,28 +61,13 @@ class ToolTypeCreate(BaseModel):
|
|||||||
data = info.data
|
data = info.data
|
||||||
if data.get("definition_type") != "compose":
|
if data.get("definition_type") != "compose":
|
||||||
return v
|
return v
|
||||||
|
|
||||||
if v is None:
|
if v is None or not v.strip():
|
||||||
raise ValueError("compose_template is required when definition_type is 'compose'")
|
raise ValueError(
|
||||||
|
"compose_template is required when definition_type is 'compose'"
|
||||||
# Replace template variables with dummy values before YAML validation
|
)
|
||||||
# to avoid YAML parsing errors with {{VAR}} syntax
|
|
||||||
sanitized = _sanitize_template_vars(v)
|
validate_compose_yaml(v)
|
||||||
|
|
||||||
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 v
|
return v
|
||||||
|
|
||||||
@field_validator("dockerfile_template")
|
@field_validator("dockerfile_template")
|
||||||
@@ -97,13 +76,15 @@ class ToolTypeCreate(BaseModel):
|
|||||||
data = info.data
|
data = info.data
|
||||||
if data.get("definition_type") != "dockerfile":
|
if data.get("definition_type") != "dockerfile":
|
||||||
return v
|
return v
|
||||||
|
|
||||||
if v is None:
|
if v is None or not v.strip():
|
||||||
raise ValueError("dockerfile_template is required when definition_type is 'dockerfile'")
|
raise ValueError(
|
||||||
|
"dockerfile_template is required when definition_type is 'dockerfile'"
|
||||||
|
)
|
||||||
|
|
||||||
if not v.strip().startswith("FROM"):
|
if not v.strip().startswith("FROM"):
|
||||||
raise ValueError("Dockerfile must start with a FROM instruction")
|
raise ValueError("Dockerfile must start with a FROM instruction")
|
||||||
|
|
||||||
return v
|
return v
|
||||||
|
|
||||||
@field_validator("interface_type")
|
@field_validator("interface_type")
|
||||||
@@ -129,57 +110,62 @@ class ToolTypeCreate(BaseModel):
|
|||||||
def validate_required_variables(cls, v: list[str], info) -> list[str]:
|
def validate_required_variables(cls, v: list[str], info) -> list[str]:
|
||||||
if not v:
|
if not v:
|
||||||
return v
|
return v
|
||||||
|
|
||||||
data = info.data
|
data = info.data
|
||||||
if data.get("definition_type") != "compose":
|
if data.get("definition_type") != "compose":
|
||||||
return v
|
return v
|
||||||
|
|
||||||
template = data.get("compose_template")
|
template = data.get("compose_template")
|
||||||
if not template:
|
if not template:
|
||||||
return v
|
return v
|
||||||
|
|
||||||
for var in v:
|
for var in v:
|
||||||
placeholder = f"{{{{{var}}}}}"
|
placeholder = f"{{{{{var}}}}}"
|
||||||
if placeholder not in template:
|
if placeholder not in template:
|
||||||
raise ValueError(f"Required variable '{var}' not found in compose template")
|
raise ValueError(
|
||||||
|
f"Required variable '{var}' not found in compose template"
|
||||||
|
)
|
||||||
|
|
||||||
return v
|
return v
|
||||||
|
|
||||||
@model_validator(mode="after")
|
@model_validator(mode="after")
|
||||||
def validate_templates(self) -> "ToolTypeCreate":
|
def validate_templates(self) -> "ToolTypeCreate":
|
||||||
if self.definition_type == "dockerfile" and self.dockerfile_template is None:
|
if self.definition_type == "manifest":
|
||||||
raise ValueError("dockerfile_template is required when definition_type is 'dockerfile'")
|
if self.manifest_id is None:
|
||||||
if self.definition_type == "compose" and self.compose_template is None:
|
raise ValueError(
|
||||||
raise ValueError("compose_template is required when definition_type is 'compose'")
|
"manifest_id is required when definition_type is 'manifest'"
|
||||||
|
)
|
||||||
|
return self
|
||||||
|
|
||||||
|
if self.definition_type == "dockerfile" and (
|
||||||
|
self.dockerfile_template is None or not self.dockerfile_template.strip()
|
||||||
|
):
|
||||||
|
raise ValueError(
|
||||||
|
"dockerfile_template is required when definition_type is 'dockerfile'"
|
||||||
|
)
|
||||||
|
if self.definition_type == "compose" and (
|
||||||
|
self.compose_template is None or not self.compose_template.strip()
|
||||||
|
):
|
||||||
|
raise ValueError(
|
||||||
|
"compose_template is required when definition_type is 'compose'"
|
||||||
|
)
|
||||||
|
|
||||||
# Validate that default_port is exposed in compose template (only if requires_port)
|
# 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:
|
if (
|
||||||
|
self.requires_port
|
||||||
|
and self.definition_type == "compose"
|
||||||
|
and self.compose_template
|
||||||
|
):
|
||||||
try:
|
try:
|
||||||
sanitized = _sanitize_template_vars(self.compose_template)
|
parsed = validate_compose_yaml(self.compose_template)
|
||||||
parsed = yaml.safe_load(sanitized)
|
except ValueError:
|
||||||
except yaml.YAMLError:
|
|
||||||
return self
|
return self
|
||||||
|
|
||||||
port_str = str(self.default_port)
|
if not check_port_exposed(parsed, self.default_port):
|
||||||
port_exposed = False
|
raise ValueError(
|
||||||
|
f"Port {self.default_port} is not exposed in the compose template. Add it to the 'ports' section."
|
||||||
if isinstance(parsed, dict) and "services" in parsed:
|
)
|
||||||
for service_name, service_config in parsed["services"].items():
|
|
||||||
if isinstance(service_config, dict) and "ports" in service_config:
|
|
||||||
for port_mapping in service_config["ports"]:
|
|
||||||
if isinstance(port_mapping, str):
|
|
||||||
if port_str in port_mapping:
|
|
||||||
port_exposed = True
|
|
||||||
break
|
|
||||||
elif isinstance(port_mapping, int) and port_mapping == self.default_port:
|
|
||||||
port_exposed = True
|
|
||||||
break
|
|
||||||
if port_exposed:
|
|
||||||
break
|
|
||||||
|
|
||||||
if not port_exposed:
|
|
||||||
raise ValueError(f"Port {self.default_port} is not exposed in the compose template. Add it to the 'ports' section.")
|
|
||||||
|
|
||||||
return self
|
return self
|
||||||
|
|
||||||
|
|
||||||
@@ -188,10 +174,12 @@ class ToolTypeUpdate(BaseModel):
|
|||||||
description: str | None = None
|
description: str | None = None
|
||||||
default_port: int | None = None
|
default_port: int | None = None
|
||||||
definition_type: str | None = None
|
definition_type: str | None = None
|
||||||
|
manifest_id: uuid.UUID | None = None
|
||||||
compose_template: str | None = None
|
compose_template: str | None = None
|
||||||
dockerfile_template: str | None = None
|
dockerfile_template: str | None = None
|
||||||
build_context: dict | None = None
|
build_context: dict | None = None
|
||||||
readiness_probe: dict | None = None
|
readiness_probe: dict | None = None
|
||||||
|
startup_command: str | None = None
|
||||||
required_variables: list[str] | None = None
|
required_variables: list[str] | None = None
|
||||||
category: str | None = None
|
category: str | None = None
|
||||||
interface_type: str | None = None
|
interface_type: str | None = None
|
||||||
@@ -202,8 +190,10 @@ class ToolTypeUpdate(BaseModel):
|
|||||||
def validate_definition_type(cls, v: str | None) -> str | None:
|
def validate_definition_type(cls, v: str | None) -> str | None:
|
||||||
if v is None:
|
if v is None:
|
||||||
return v
|
return v
|
||||||
if v not in ("compose", "dockerfile"):
|
if v not in ("compose", "dockerfile", "manifest"):
|
||||||
raise ValueError("definition_type must be 'compose' or 'dockerfile'")
|
raise ValueError(
|
||||||
|
"definition_type must be 'compose', 'dockerfile', or 'manifest'"
|
||||||
|
)
|
||||||
return v
|
return v
|
||||||
|
|
||||||
@field_validator("interface_type")
|
@field_validator("interface_type")
|
||||||
@@ -220,29 +210,13 @@ class ToolTypeUpdate(BaseModel):
|
|||||||
def validate_compose_template(cls, v: str | None, info) -> str | None:
|
def validate_compose_template(cls, v: str | None, info) -> str | None:
|
||||||
if v is None:
|
if v is None:
|
||||||
return v
|
return v
|
||||||
|
|
||||||
data = info.data
|
data = info.data
|
||||||
definition_type = data.get("definition_type")
|
definition_type = data.get("definition_type")
|
||||||
if definition_type and definition_type != "compose":
|
if definition_type and definition_type != "compose":
|
||||||
return v
|
return v
|
||||||
|
|
||||||
# Replace template variables with dummy values before YAML validation
|
validate_compose_yaml(v)
|
||||||
sanitized = _sanitize_template_vars(v)
|
|
||||||
|
|
||||||
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 v
|
return v
|
||||||
|
|
||||||
@field_validator("dockerfile_template")
|
@field_validator("dockerfile_template")
|
||||||
@@ -250,15 +224,15 @@ class ToolTypeUpdate(BaseModel):
|
|||||||
def validate_dockerfile_template(cls, v: str | None, info) -> str | None:
|
def validate_dockerfile_template(cls, v: str | None, info) -> str | None:
|
||||||
if v is None:
|
if v is None:
|
||||||
return v
|
return v
|
||||||
|
|
||||||
data = info.data
|
data = info.data
|
||||||
definition_type = data.get("definition_type")
|
definition_type = data.get("definition_type")
|
||||||
if definition_type and definition_type != "dockerfile":
|
if definition_type and definition_type != "dockerfile":
|
||||||
return v
|
return v
|
||||||
|
|
||||||
if not v.strip().startswith("FROM"):
|
if not v.strip().startswith("FROM"):
|
||||||
raise ValueError("Dockerfile must start with a FROM instruction")
|
raise ValueError("Dockerfile must start with a FROM instruction")
|
||||||
|
|
||||||
return v
|
return v
|
||||||
|
|
||||||
|
|
||||||
@@ -274,10 +248,12 @@ class ToolTypeResponse(BaseModel):
|
|||||||
requires_port: bool
|
requires_port: bool
|
||||||
default_port: int
|
default_port: int
|
||||||
definition_type: str
|
definition_type: str
|
||||||
|
manifest_id: uuid.UUID | None
|
||||||
compose_template: str | None
|
compose_template: str | None
|
||||||
dockerfile_template: str | None
|
dockerfile_template: str | None
|
||||||
build_context: dict | None
|
build_context: dict | None
|
||||||
readiness_probe: dict | None
|
readiness_probe: dict | None
|
||||||
|
startup_command: str | None
|
||||||
required_variables: list[str]
|
required_variables: list[str]
|
||||||
created_by_id: uuid.UUID | None
|
created_by_id: uuid.UUID | None
|
||||||
created_at: datetime
|
created_at: datetime
|
||||||
@@ -308,22 +284,27 @@ async def create_tool_type(
|
|||||||
"""
|
"""
|
||||||
user = await _get_user(session, user_id)
|
user = await _get_user(session, user_id)
|
||||||
await _require_admin(user)
|
await _require_admin(user)
|
||||||
|
|
||||||
# Check for duplicate name
|
# Check for duplicate name
|
||||||
existing = await session.scalar(select(ToolType).where(ToolType.name == data.name))
|
existing = await session.scalar(select(ToolType).where(ToolType.name == data.name))
|
||||||
if existing:
|
if existing:
|
||||||
raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail="tool type with this name already exists")
|
raise HTTPException(
|
||||||
|
status_code=status.HTTP_409_CONFLICT,
|
||||||
|
detail="tool type with this name already exists",
|
||||||
|
)
|
||||||
|
|
||||||
tool_type = ToolType(
|
tool_type = ToolType(
|
||||||
name=data.name,
|
name=data.name,
|
||||||
display_name=data.display_name,
|
display_name=data.display_name,
|
||||||
description=data.description,
|
description=data.description,
|
||||||
default_port=data.default_port,
|
default_port=data.default_port,
|
||||||
definition_type=data.definition_type,
|
definition_type=data.definition_type,
|
||||||
|
manifest_id=data.manifest_id,
|
||||||
compose_template=data.compose_template,
|
compose_template=data.compose_template,
|
||||||
dockerfile_template=data.dockerfile_template,
|
dockerfile_template=data.dockerfile_template,
|
||||||
build_context=data.build_context,
|
build_context=data.build_context,
|
||||||
readiness_probe=data.readiness_probe,
|
readiness_probe=data.readiness_probe,
|
||||||
|
startup_command=data.startup_command,
|
||||||
required_variables=data.required_variables,
|
required_variables=data.required_variables,
|
||||||
category=data.category,
|
category=data.category,
|
||||||
interface_type=data.interface_type,
|
interface_type=data.interface_type,
|
||||||
@@ -384,7 +365,9 @@ async def get_tool_type(
|
|||||||
await _get_user(session, user_id)
|
await _get_user(session, user_id)
|
||||||
tool_type = await session.get(ToolType, tool_type_id)
|
tool_type = await session.get(ToolType, tool_type_id)
|
||||||
if tool_type is None:
|
if tool_type is None:
|
||||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="tool type not found")
|
raise HTTPException(
|
||||||
|
status_code=status.HTTP_404_NOT_FOUND, detail="tool type not found"
|
||||||
|
)
|
||||||
return tool_type
|
return tool_type
|
||||||
|
|
||||||
|
|
||||||
@@ -413,15 +396,17 @@ async def update_tool_type(
|
|||||||
"""
|
"""
|
||||||
user = await _get_user(session, user_id)
|
user = await _get_user(session, user_id)
|
||||||
await _require_admin(user)
|
await _require_admin(user)
|
||||||
|
|
||||||
tool_type = await session.get(ToolType, tool_type_id)
|
tool_type = await session.get(ToolType, tool_type_id)
|
||||||
if tool_type is None:
|
if tool_type is None:
|
||||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="tool type not found")
|
raise HTTPException(
|
||||||
|
status_code=status.HTTP_404_NOT_FOUND, detail="tool type not found"
|
||||||
|
)
|
||||||
|
|
||||||
# Built-in tool types can now be modified
|
# Built-in tool types can now be modified
|
||||||
|
|
||||||
update_data = data.model_dump(exclude_unset=True)
|
update_data = data.model_dump(exclude_unset=True)
|
||||||
|
|
||||||
# Validate port if being updated
|
# Validate port if being updated
|
||||||
requires_port = update_data.get("requires_port", tool_type.requires_port)
|
requires_port = update_data.get("requires_port", tool_type.requires_port)
|
||||||
if "default_port" in update_data and requires_port:
|
if "default_port" in update_data and requires_port:
|
||||||
@@ -429,67 +414,48 @@ async def update_tool_type(
|
|||||||
if new_port <= 0 or new_port > 65535:
|
if new_port <= 0 or new_port > 65535:
|
||||||
raise HTTPException(
|
raise HTTPException(
|
||||||
status_code=status.HTTP_400_BAD_REQUEST,
|
status_code=status.HTTP_400_BAD_REQUEST,
|
||||||
detail="Port must be between 1 and 65535"
|
detail="Port must be between 1 and 65535",
|
||||||
)
|
)
|
||||||
|
|
||||||
# Only validate port exposure for compose definitions
|
# Only validate port exposure for compose definitions
|
||||||
definition_type = update_data.get("definition_type", tool_type.definition_type)
|
definition_type = update_data.get("definition_type", tool_type.definition_type)
|
||||||
if definition_type == "compose":
|
if definition_type == "compose":
|
||||||
template = update_data.get("compose_template", tool_type.compose_template)
|
template = update_data.get("compose_template", tool_type.compose_template)
|
||||||
if template:
|
if template:
|
||||||
try:
|
try:
|
||||||
sanitized = _sanitize_template_vars(template)
|
parsed = validate_compose_yaml(template)
|
||||||
parsed = yaml.safe_load(sanitized)
|
if not check_port_exposed(parsed, new_port):
|
||||||
except yaml.YAMLError:
|
|
||||||
parsed = None
|
|
||||||
|
|
||||||
if parsed and isinstance(parsed, dict) and "services" in parsed:
|
|
||||||
port_str = str(new_port)
|
|
||||||
port_exposed = 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:
|
|
||||||
port_exposed = True
|
|
||||||
break
|
|
||||||
elif isinstance(port_mapping, int) and port_mapping == new_port:
|
|
||||||
port_exposed = True
|
|
||||||
break
|
|
||||||
if port_exposed:
|
|
||||||
break
|
|
||||||
|
|
||||||
if not port_exposed:
|
|
||||||
raise HTTPException(
|
raise HTTPException(
|
||||||
status_code=status.HTTP_400_BAD_REQUEST,
|
status_code=status.HTTP_400_BAD_REQUEST,
|
||||||
detail=f"Port {new_port} is not exposed in the compose template"
|
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
|
# Validate required variables for compose definitions
|
||||||
definition_type = update_data.get("definition_type", tool_type.definition_type)
|
definition_type = update_data.get("definition_type", tool_type.definition_type)
|
||||||
if definition_type == "compose":
|
if definition_type == "compose":
|
||||||
if "required_variables" in update_data and "compose_template" in update_data:
|
if "required_variables" in update_data and "compose_template" in update_data:
|
||||||
template = update_data["compose_template"]
|
validate_required_variables(
|
||||||
for var in update_data["required_variables"]:
|
update_data["compose_template"], update_data["required_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"
|
|
||||||
)
|
|
||||||
elif "required_variables" in update_data:
|
elif "required_variables" in update_data:
|
||||||
template = tool_type.compose_template
|
template = tool_type.compose_template
|
||||||
if template:
|
if template:
|
||||||
for var in update_data["required_variables"]:
|
validate_required_variables(template, update_data["required_variables"])
|
||||||
placeholder = f"{{{{{var}}}}}"
|
|
||||||
if placeholder not in template:
|
# When switching to manifest, clear legacy templates
|
||||||
raise HTTPException(
|
if definition_type == "manifest":
|
||||||
status_code=status.HTTP_400_BAD_REQUEST,
|
if "manifest_id" in update_data:
|
||||||
detail=f"Required variable '{var}' not found in compose template"
|
tool_type.manifest_id = update_data["manifest_id"]
|
||||||
)
|
tool_type.compose_template = None
|
||||||
|
tool_type.dockerfile_template = None
|
||||||
|
|
||||||
for field, value in update_data.items():
|
for field, value in update_data.items():
|
||||||
setattr(tool_type, field, value)
|
setattr(tool_type, field, value)
|
||||||
|
|
||||||
await session.commit()
|
await session.commit()
|
||||||
await session.refresh(tool_type)
|
await session.refresh(tool_type)
|
||||||
return tool_type
|
return tool_type
|
||||||
@@ -530,16 +496,9 @@ async def validate_tool_type_template(
|
|||||||
errors.append("Compose template is required")
|
errors.append("Compose template is required")
|
||||||
else:
|
else:
|
||||||
try:
|
try:
|
||||||
sanitized = _sanitize_template_vars(data.compose_template)
|
validate_compose_yaml(data.compose_template)
|
||||||
parsed = yaml.safe_load(sanitized)
|
except ValueError as e:
|
||||||
if not isinstance(parsed, dict):
|
errors.append(str(e))
|
||||||
errors.append("Compose template must be a YAML mapping")
|
|
||||||
elif "services" not in parsed:
|
|
||||||
errors.append("Compose template must contain 'services' key")
|
|
||||||
elif not parsed["services"]:
|
|
||||||
errors.append("Compose template must define at least one service")
|
|
||||||
except yaml.YAMLError as e:
|
|
||||||
errors.append(f"Invalid YAML: {e}")
|
|
||||||
|
|
||||||
elif data.definition_type == "dockerfile":
|
elif data.definition_type == "dockerfile":
|
||||||
if not data.dockerfile_template:
|
if not data.dockerfile_template:
|
||||||
@@ -547,8 +506,11 @@ async def validate_tool_type_template(
|
|||||||
elif not data.dockerfile_template.strip().startswith("FROM"):
|
elif not data.dockerfile_template.strip().startswith("FROM"):
|
||||||
errors.append("Dockerfile must start with a FROM instruction")
|
errors.append("Dockerfile must start with a FROM instruction")
|
||||||
|
|
||||||
|
elif data.definition_type == "manifest":
|
||||||
|
pass # Manifest validation is handled separately
|
||||||
|
|
||||||
else:
|
else:
|
||||||
errors.append("definition_type must be 'compose' or 'dockerfile'")
|
errors.append("definition_type must be 'compose', 'dockerfile', or 'manifest'")
|
||||||
|
|
||||||
return {
|
return {
|
||||||
"valid": len(errors) == 0,
|
"valid": len(errors) == 0,
|
||||||
@@ -579,32 +541,31 @@ async def validate_tool_type(
|
|||||||
await _get_user(session, user_id)
|
await _get_user(session, user_id)
|
||||||
tool_type = await session.get(ToolType, tool_type_id)
|
tool_type = await session.get(ToolType, tool_type_id)
|
||||||
if tool_type is None:
|
if tool_type is None:
|
||||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="tool type not found")
|
raise HTTPException(
|
||||||
|
status_code=status.HTTP_404_NOT_FOUND, detail="tool type not found"
|
||||||
|
)
|
||||||
|
|
||||||
errors = []
|
errors = []
|
||||||
|
|
||||||
if tool_type.definition_type == "compose":
|
if tool_type.definition_type == "compose":
|
||||||
if not tool_type.compose_template:
|
if not tool_type.compose_template:
|
||||||
errors.append("Compose template is empty")
|
errors.append("Compose template is empty")
|
||||||
else:
|
else:
|
||||||
try:
|
try:
|
||||||
sanitized = _sanitize_template_vars(tool_type.compose_template)
|
validate_compose_yaml(tool_type.compose_template)
|
||||||
parsed = yaml.safe_load(sanitized)
|
except ValueError as e:
|
||||||
if not isinstance(parsed, dict):
|
errors.append(str(e))
|
||||||
errors.append("Compose template must be a YAML mapping")
|
|
||||||
elif "services" not in parsed:
|
|
||||||
errors.append("Compose template must contain 'services' key")
|
|
||||||
elif not parsed["services"]:
|
|
||||||
errors.append("Compose template must define at least one service")
|
|
||||||
except yaml.YAMLError as e:
|
|
||||||
errors.append(f"Invalid YAML: {e}")
|
|
||||||
|
|
||||||
elif tool_type.definition_type == "dockerfile":
|
elif tool_type.definition_type == "dockerfile":
|
||||||
if not tool_type.dockerfile_template:
|
if not tool_type.dockerfile_template:
|
||||||
errors.append("Dockerfile template is empty")
|
errors.append("Dockerfile template is empty")
|
||||||
elif not tool_type.dockerfile_template.strip().startswith("FROM"):
|
elif not tool_type.dockerfile_template.strip().startswith("FROM"):
|
||||||
errors.append("Dockerfile must start with a FROM instruction")
|
errors.append("Dockerfile must start with a FROM instruction")
|
||||||
|
|
||||||
|
elif tool_type.definition_type == "manifest":
|
||||||
|
if not tool_type.manifest_id:
|
||||||
|
errors.append("Manifest reference is missing")
|
||||||
|
|
||||||
return {
|
return {
|
||||||
"valid": len(errors) == 0,
|
"valid": len(errors) == 0,
|
||||||
"errors": errors,
|
"errors": errors,
|
||||||
@@ -634,12 +595,14 @@ async def delete_tool_type(
|
|||||||
"""
|
"""
|
||||||
user = await _get_user(session, user_id)
|
user = await _get_user(session, user_id)
|
||||||
await _require_admin(user)
|
await _require_admin(user)
|
||||||
|
|
||||||
tool_type = await session.get(ToolType, tool_type_id)
|
tool_type = await session.get(ToolType, tool_type_id)
|
||||||
if tool_type is None:
|
if tool_type is None:
|
||||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="tool type not found")
|
raise HTTPException(
|
||||||
|
status_code=status.HTTP_404_NOT_FOUND, detail="tool type not found"
|
||||||
|
)
|
||||||
|
|
||||||
# Built-in tool types can now be deleted
|
# Built-in tool types can now be deleted
|
||||||
|
|
||||||
await session.delete(tool_type)
|
await session.delete(tool_type)
|
||||||
await session.commit()
|
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,29 +1,22 @@
|
|||||||
import logging
|
import logging
|
||||||
import uuid
|
import uuid
|
||||||
|
|
||||||
from fastapi import APIRouter, Depends, HTTPException, status
|
from fastapi import APIRouter, Depends
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
|
||||||
from pydantic import BaseModel, ConfigDict
|
from pydantic import BaseModel, ConfigDict
|
||||||
from sqlalchemy import select
|
from sqlalchemy import select
|
||||||
from sqlalchemy.ext.asyncio import AsyncSession
|
from sqlalchemy.ext.asyncio import AsyncSession
|
||||||
|
|
||||||
from src.auth.dependencies import get_current_user_id, get_db_session
|
from src.auth.dependencies import _get_user, get_current_user_id, get_db_session
|
||||||
from src.models.user import User
|
|
||||||
from src.models.user_config import UserConfig
|
from src.models.user_config import UserConfig
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
router = APIRouter(prefix="/users/me", tags=["user-config"])
|
router = APIRouter(prefix="/users/me", tags=["user-config"])
|
||||||
|
|
||||||
|
|
||||||
async def _get_user(session: AsyncSession, user_id: uuid.UUID) -> User:
|
async def _get_or_create_config(
|
||||||
"""Fetch a user by ID or raise 401 if not found."""
|
session: AsyncSession, user_id: uuid.UUID
|
||||||
user = await session.get(User, user_id)
|
) -> UserConfig:
|
||||||
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.
|
"""Get or create user config record.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
@@ -33,7 +26,9 @@ async def _get_or_create_config(session: AsyncSession, user_id: uuid.UUID) -> Us
|
|||||||
Returns:
|
Returns:
|
||||||
The user's config, creating a new one if it doesn't exist.
|
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))
|
result = await session.execute(
|
||||||
|
select(UserConfig).where(UserConfig.user_id == user_id)
|
||||||
|
)
|
||||||
config = result.scalar_one_or_none()
|
config = result.scalar_one_or_none()
|
||||||
if config is None:
|
if config is None:
|
||||||
config = UserConfig(user_id=user_id, config={})
|
config = UserConfig(user_id=user_id, config={})
|
||||||
@@ -51,6 +46,8 @@ class UserConfigResponse(BaseModel):
|
|||||||
git_user_name: str | None = None
|
git_user_name: str | None = None
|
||||||
git_user_email: str | None = None
|
git_user_email: str | None = None
|
||||||
last_session_id: str | None = None
|
last_session_id: str | None = None
|
||||||
|
notification_mute_categories: list[str] | None = None
|
||||||
|
notification_toast_level: str | None = None
|
||||||
|
|
||||||
|
|
||||||
class UserConfigUpdate(BaseModel):
|
class UserConfigUpdate(BaseModel):
|
||||||
@@ -59,6 +56,8 @@ class UserConfigUpdate(BaseModel):
|
|||||||
git_user_name: str | None = None
|
git_user_name: str | None = None
|
||||||
git_user_email: str | None = None
|
git_user_email: str | None = None
|
||||||
last_session_id: str | None = None
|
last_session_id: str | None = None
|
||||||
|
notification_mute_categories: list[str] | None = None
|
||||||
|
notification_toast_level: str | None = None
|
||||||
|
|
||||||
|
|
||||||
@router.get(
|
@router.get(
|
||||||
@@ -111,11 +110,11 @@ async def update_user_config(
|
|||||||
|
|
||||||
# Merge updates
|
# Merge updates
|
||||||
update_data = data.model_dump(exclude_unset=True)
|
update_data = data.model_dump(exclude_unset=True)
|
||||||
logger.info("Updating user config for user %s: %s", user_id, update_data)
|
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
|
# SQLAlchemy JSON doesn't track dict mutations, so we replace the whole dict
|
||||||
config.config = {**config.config, **update_data}
|
config.config = {**config.config, **update_data}
|
||||||
|
|
||||||
await session.commit()
|
await session.commit()
|
||||||
await session.refresh(config)
|
await session.refresh(config)
|
||||||
logger.info("Updated config: %s", config.config)
|
logger.debug("Updated config: %s", config.config)
|
||||||
return UserConfigResponse.model_validate(config.config)
|
return UserConfigResponse.model_validate(config.config)
|
||||||
|
|||||||
@@ -5,7 +5,7 @@ from fastapi import APIRouter, Depends, HTTPException, UploadFile, status
|
|||||||
from pydantic import BaseModel, ConfigDict
|
from pydantic import BaseModel, ConfigDict
|
||||||
from sqlalchemy.ext.asyncio import AsyncSession
|
from sqlalchemy.ext.asyncio import AsyncSession
|
||||||
|
|
||||||
from src.auth.dependencies import get_current_user_id, get_db_session
|
from src.auth.dependencies import _get_user, get_current_user_id, get_db_session
|
||||||
from src.models.user import User
|
from src.models.user import User
|
||||||
|
|
||||||
router = APIRouter(prefix="/users", tags=["users"])
|
router = APIRouter(prefix="/users", tags=["users"])
|
||||||
@@ -16,14 +16,6 @@ ALLOWED_CONTENT_TYPES = {"image/png", "image/jpeg", "image/jpg"}
|
|||||||
MAX_AVATAR_SIZE = 2 * 1024 * 1024 # 2MB
|
MAX_AVATAR_SIZE = 2 * 1024 * 1024 # 2MB
|
||||||
|
|
||||||
|
|
||||||
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
|
|
||||||
|
|
||||||
|
|
||||||
class UserProfileResponse(BaseModel):
|
class UserProfileResponse(BaseModel):
|
||||||
model_config = ConfigDict(from_attributes=True)
|
model_config = ConfigDict(from_attributes=True)
|
||||||
|
|
||||||
|
|||||||
@@ -7,6 +7,7 @@ from sqlalchemy.ext.asyncio import AsyncSession
|
|||||||
from src.auth.session import decode_session_cookie
|
from src.auth.session import decode_session_cookie
|
||||||
from src.config import Settings
|
from src.config import Settings
|
||||||
from src.database import SessionLocal
|
from src.database import SessionLocal
|
||||||
|
from src.models.project import Project
|
||||||
from src.models.user import User
|
from src.models.user import User
|
||||||
|
|
||||||
|
|
||||||
@@ -47,3 +48,39 @@ async def get_current_user(
|
|||||||
if user is None:
|
if user is None:
|
||||||
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="user not found")
|
raise HTTPException(status_code=status.HTTP_401_UNAUTHORIZED, detail="user not found")
|
||||||
return user
|
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,15 +1,52 @@
|
|||||||
|
"""Structured JSON logging configuration."""
|
||||||
|
|
||||||
|
import json
|
||||||
import logging
|
import logging
|
||||||
import sys
|
import sys
|
||||||
import time
|
import time
|
||||||
import traceback
|
import traceback
|
||||||
from typing import Callable
|
from collections.abc import Callable
|
||||||
|
|
||||||
from fastapi import Request, Response
|
from fastapi import Request, Response
|
||||||
from starlette.middleware.base import BaseHTTPMiddleware
|
from starlette.middleware.base import BaseHTTPMiddleware
|
||||||
|
|
||||||
|
from src.services.correlation import get_correlation_id
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class CorrelationIdFilter(logging.Filter):
|
||||||
|
"""Inject correlation_id into every log record from context var."""
|
||||||
|
|
||||||
|
def filter(self, record: logging.LogRecord) -> bool:
|
||||||
|
record.correlation_id = get_correlation_id() # type: ignore[attr-defined]
|
||||||
|
return True
|
||||||
|
|
||||||
|
|
||||||
|
class JSONFormatter(logging.Formatter):
|
||||||
|
"""Emit log records as single-line JSON."""
|
||||||
|
|
||||||
|
def format(self, record: logging.LogRecord) -> str:
|
||||||
|
log_obj: dict = {
|
||||||
|
"timestamp": self.formatTime(record),
|
||||||
|
"level": record.levelname,
|
||||||
|
"logger": record.name,
|
||||||
|
"message": record.getMessage(),
|
||||||
|
"correlation_id": getattr(record, "correlation_id", None),
|
||||||
|
}
|
||||||
|
# Optional extra fields
|
||||||
|
for key in ("instance_id", "event_type"):
|
||||||
|
value = getattr(record, key, None)
|
||||||
|
if value is not None:
|
||||||
|
log_obj[key] = value
|
||||||
|
if record.exc_info:
|
||||||
|
log_obj["exception"] = self.formatException(record.exc_info)
|
||||||
|
return json.dumps(log_obj, default=str)
|
||||||
|
|
||||||
|
def formatTime(self, record: logging.LogRecord, datefmt: str | None = None) -> str:
|
||||||
|
return time.strftime("%Y-%m-%dT%H:%M:%S", time.gmtime(record.created))
|
||||||
|
|
||||||
|
|
||||||
class RequestLoggingMiddleware(BaseHTTPMiddleware):
|
class RequestLoggingMiddleware(BaseHTTPMiddleware):
|
||||||
"""Log all HTTP requests with timing and status codes."""
|
"""Log all HTTP requests with timing and status codes."""
|
||||||
|
|
||||||
@@ -17,7 +54,6 @@ class RequestLoggingMiddleware(BaseHTTPMiddleware):
|
|||||||
start_time = time.time()
|
start_time = time.time()
|
||||||
client_host = request.client.host if request.client else "unknown"
|
client_host = request.client.host if request.client else "unknown"
|
||||||
|
|
||||||
# Log the incoming request
|
|
||||||
logger.info(
|
logger.info(
|
||||||
"→ Request: %s %s (client: %s)",
|
"→ Request: %s %s (client: %s)",
|
||||||
request.method,
|
request.method,
|
||||||
@@ -29,7 +65,6 @@ class RequestLoggingMiddleware(BaseHTTPMiddleware):
|
|||||||
response = await call_next(request)
|
response = await call_next(request)
|
||||||
duration = time.time() - start_time
|
duration = time.time() - start_time
|
||||||
|
|
||||||
# Log the response
|
|
||||||
logger.info(
|
logger.info(
|
||||||
"← Response: %s %s → %d (%dms)",
|
"← Response: %s %s → %d (%dms)",
|
||||||
request.method,
|
request.method,
|
||||||
@@ -69,15 +104,13 @@ class ExceptionLoggingMiddleware(BaseHTTPMiddleware):
|
|||||||
|
|
||||||
|
|
||||||
def configure_logging(level: int = logging.INFO) -> None:
|
def configure_logging(level: int = logging.INFO) -> None:
|
||||||
"""Configure structured logging for the application."""
|
"""Configure structured JSON logging for the application."""
|
||||||
formatter = logging.Formatter(
|
formatter = JSONFormatter()
|
||||||
fmt="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
|
|
||||||
datefmt="%Y-%m-%d %H:%M:%S",
|
|
||||||
)
|
|
||||||
|
|
||||||
# Console handler
|
# Console handler
|
||||||
console_handler = logging.StreamHandler(sys.stdout)
|
console_handler = logging.StreamHandler(sys.stdout)
|
||||||
console_handler.setFormatter(formatter)
|
console_handler.setFormatter(formatter)
|
||||||
|
console_handler.addFilter(CorrelationIdFilter())
|
||||||
|
|
||||||
# Configure root logger
|
# Configure root logger
|
||||||
root_logger = logging.getLogger()
|
root_logger = logging.getLogger()
|
||||||
|
|||||||
+34
-7
@@ -1,4 +1,3 @@
|
|||||||
import json
|
|
||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
|
|
||||||
@@ -7,31 +6,36 @@ from fastapi.exceptions import RequestValidationError
|
|||||||
from fastapi.middleware.cors import CORSMiddleware
|
from fastapi.middleware.cors import CORSMiddleware
|
||||||
from fastapi.responses import JSONResponse
|
from fastapi.responses import JSONResponse
|
||||||
from fastapi.staticfiles import StaticFiles
|
from fastapi.staticfiles import StaticFiles
|
||||||
from sqlalchemy import text
|
|
||||||
|
|
||||||
from src.api.auth import router as auth_router
|
from src.api.auth import router as auth_router
|
||||||
from src.api.dashboard import router as dashboard_router
|
from src.api.dashboard import router as dashboard_router
|
||||||
|
from src.api.events import router as events_router
|
||||||
from src.api.git_repositories import router as git_repositories_router
|
from src.api.git_repositories import router as git_repositories_router
|
||||||
from src.api.health import router as health_router
|
from src.api.health import router as health_router
|
||||||
from src.api.projects import router as projects_router
|
from src.api.projects import router as projects_router
|
||||||
from src.api.ssh_keys import router as ssh_keys_router
|
from src.api.ssh_keys import router as ssh_keys_router
|
||||||
from src.api.terminal import router as terminal_router
|
from src.api.terminal import router as terminal_router
|
||||||
from src.api.instance_proxy import router as instance_proxy_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.config_profiles import router as config_profiles_router
|
||||||
from src.api.tool_configs import router as tool_configs_router
|
from src.api.tool_definitions import router as tool_definitions_router
|
||||||
from src.api.tool_instances import router as tool_instances_router
|
from src.api.tool_instances import router as tool_instances_router
|
||||||
from src.api.tool_instances import sessions_router
|
from src.api.tool_instances import sessions_router
|
||||||
from src.api.tool_types import router as tool_types_router
|
from src.api.tool_types import router as tool_types_router
|
||||||
|
from src.api.notifications import router as notifications_router
|
||||||
from src.api.user_config import router as user_config_router
|
from src.api.user_config import router as user_config_router
|
||||||
from src.api.users import router as users_router
|
from src.api.users import router as users_router
|
||||||
from src.config import Settings
|
from src.config import Settings
|
||||||
|
from src.models.notification import Notification # noqa: F401 – Alembic model discovery
|
||||||
|
from src.models.terminal_session import TerminalSessionModel # noqa: F401 – Alembic model discovery
|
||||||
from src.database import init_database
|
from src.database import init_database
|
||||||
from src.logging_config import (
|
from src.logging_config import (
|
||||||
ExceptionLoggingMiddleware,
|
ExceptionLoggingMiddleware,
|
||||||
RequestLoggingMiddleware,
|
RequestLoggingMiddleware,
|
||||||
configure_logging,
|
configure_logging,
|
||||||
)
|
)
|
||||||
|
from src.services.correlation import CorrelationIdMiddleware
|
||||||
|
from src.services.event_bus import InstanceEventBus
|
||||||
|
from src.services.health_monitor import HealthMonitor
|
||||||
|
|
||||||
# Configure logging early
|
# Configure logging early
|
||||||
log_level = os.getenv("LOG_LEVEL", "INFO").upper()
|
log_level = os.getenv("LOG_LEVEL", "INFO").upper()
|
||||||
@@ -56,6 +60,7 @@ app.add_middleware(
|
|||||||
allow_headers=["*"],
|
allow_headers=["*"],
|
||||||
)
|
)
|
||||||
|
|
||||||
|
app.add_middleware(CorrelationIdMiddleware)
|
||||||
app.add_middleware(RequestLoggingMiddleware)
|
app.add_middleware(RequestLoggingMiddleware)
|
||||||
app.add_middleware(ExceptionLoggingMiddleware)
|
app.add_middleware(ExceptionLoggingMiddleware)
|
||||||
|
|
||||||
@@ -68,7 +73,9 @@ def _sanitize_validation_errors(errors):
|
|||||||
"type": error.get("type"),
|
"type": error.get("type"),
|
||||||
"loc": error.get("loc"),
|
"loc": error.get("loc"),
|
||||||
"msg": error.get("msg"),
|
"msg": error.get("msg"),
|
||||||
"input": str(error.get("input")) if error.get("input") is not None else None,
|
"input": str(error.get("input"))
|
||||||
|
if error.get("input") is not None
|
||||||
|
else None,
|
||||||
}
|
}
|
||||||
# Convert ctx to safe format
|
# Convert ctx to safe format
|
||||||
ctx = error.get("ctx")
|
ctx = error.get("ctx")
|
||||||
@@ -103,6 +110,11 @@ async def validation_exception_handler(request: Request, exc: RequestValidationE
|
|||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# Global services
|
||||||
|
_event_bus = InstanceEventBus()
|
||||||
|
_health_monitor = HealthMonitor(_event_bus)
|
||||||
|
|
||||||
|
|
||||||
@app.on_event("startup")
|
@app.on_event("startup")
|
||||||
async def on_startup():
|
async def on_startup():
|
||||||
logger.info("Starting up Headquarter API...")
|
logger.info("Starting up Headquarter API...")
|
||||||
@@ -112,10 +124,24 @@ async def on_startup():
|
|||||||
if not db_ready:
|
if not db_ready:
|
||||||
logger.error("Database initialization failed. Shutting down.")
|
logger.error("Database initialization failed. Shutting down.")
|
||||||
import sys
|
import sys
|
||||||
|
|
||||||
sys.exit(1)
|
sys.exit(1)
|
||||||
|
|
||||||
|
# Start background health monitor
|
||||||
|
_health_monitor.start()
|
||||||
|
logger.info("Health monitor started")
|
||||||
|
|
||||||
logger.info("Startup complete.")
|
logger.info("Startup complete.")
|
||||||
|
|
||||||
|
|
||||||
|
@app.on_event("shutdown")
|
||||||
|
async def on_shutdown():
|
||||||
|
logger.info("Shutting down Headquarter API...")
|
||||||
|
_health_monitor.stop()
|
||||||
|
logger.info("Health monitor stopped")
|
||||||
|
logger.info("Shutdown complete.")
|
||||||
|
|
||||||
|
|
||||||
app.include_router(health_router)
|
app.include_router(health_router)
|
||||||
app.include_router(auth_router)
|
app.include_router(auth_router)
|
||||||
app.include_router(dashboard_router)
|
app.include_router(dashboard_router)
|
||||||
@@ -125,11 +151,12 @@ app.include_router(ssh_keys_router)
|
|||||||
app.include_router(git_repositories_router)
|
app.include_router(git_repositories_router)
|
||||||
app.include_router(user_config_router)
|
app.include_router(user_config_router)
|
||||||
app.include_router(tool_types_router)
|
app.include_router(tool_types_router)
|
||||||
app.include_router(config_folders_router)
|
app.include_router(tool_definitions_router)
|
||||||
app.include_router(config_profiles_router)
|
app.include_router(config_profiles_router)
|
||||||
app.include_router(tool_instances_router)
|
app.include_router(tool_instances_router)
|
||||||
app.include_router(tool_configs_router)
|
|
||||||
app.include_router(sessions_router)
|
app.include_router(sessions_router)
|
||||||
app.include_router(instance_proxy_router)
|
app.include_router(instance_proxy_router)
|
||||||
app.include_router(terminal_router)
|
app.include_router(terminal_router)
|
||||||
|
app.include_router(events_router)
|
||||||
|
app.include_router(notifications_router)
|
||||||
app.mount("/uploads", StaticFiles(directory="uploads"), name="uploads")
|
app.mount("/uploads", StaticFiles(directory="uploads"), name="uploads")
|
||||||
|
|||||||
@@ -1,12 +1,32 @@
|
|||||||
from src.models.base import Base
|
from src.models.base import Base
|
||||||
from src.models.config_folder import ConfigFolder
|
|
||||||
from src.models.config_profile import ConfigProfile, ConfigProfileInclude
|
from src.models.config_profile import ConfigProfile, ConfigProfileInclude
|
||||||
from src.models.git_repository import GitRepository
|
from src.models.git_repository import GitRepository
|
||||||
|
from src.models.health_check import HealthCheck
|
||||||
|
from src.models.instance_event import InstanceEvent
|
||||||
|
from src.models.notification import Notification
|
||||||
from src.models.project import Project
|
from src.models.project import Project
|
||||||
from src.models.ssh_key import SSHKey
|
from src.models.ssh_key import SSHKey
|
||||||
|
from src.models.terminal_session import TerminalSessionModel
|
||||||
|
from src.models.tool_definition_manifest import ToolDefinitionManifest
|
||||||
from src.models.tool_instance import ToolInstance
|
from src.models.tool_instance import ToolInstance
|
||||||
from src.models.tool_type import ToolType
|
from src.models.tool_type import ToolType
|
||||||
from src.models.user import User
|
from src.models.user import User
|
||||||
from src.models.user_config import UserConfig
|
from src.models.user_config import UserConfig
|
||||||
|
|
||||||
__all__ = ["Base", "ConfigFolder", "ConfigProfile", "ConfigProfileInclude", "GitRepository", "Project", "SSHKey", "ToolInstance", "ToolType", "User", "UserConfig"]
|
__all__ = [
|
||||||
|
"Base",
|
||||||
|
"ConfigProfile",
|
||||||
|
"ConfigProfileInclude",
|
||||||
|
"GitRepository",
|
||||||
|
"HealthCheck",
|
||||||
|
"InstanceEvent",
|
||||||
|
"Notification",
|
||||||
|
"Project",
|
||||||
|
"SSHKey",
|
||||||
|
"TerminalSessionModel",
|
||||||
|
"ToolDefinitionManifest",
|
||||||
|
"ToolInstance",
|
||||||
|
"ToolType",
|
||||||
|
"User",
|
||||||
|
"UserConfig",
|
||||||
|
]
|
||||||
|
|||||||
@@ -1,31 +0,0 @@
|
|||||||
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()
|
|
||||||
@@ -39,6 +39,9 @@ class ConfigProfile(UUIDPrimaryKeyMixin, TimestampMixin, Base):
|
|||||||
files: Mapped[dict] = mapped_column(
|
files: Mapped[dict] = mapped_column(
|
||||||
JSON, default=dict, nullable=False
|
JSON, default=dict, nullable=False
|
||||||
) # {"rel/path": "content", ...}
|
) # {"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)
|
is_default: Mapped[bool] = mapped_column(Boolean, default=False, nullable=False)
|
||||||
|
|
||||||
user: Mapped["User"] = relationship()
|
user: Mapped["User"] = relationship()
|
||||||
|
|||||||
@@ -19,7 +19,7 @@ class GitRepository(UUIDPrimaryKeyMixin, TimestampMixin, Base):
|
|||||||
|
|
||||||
name: Mapped[str] = mapped_column(String(255))
|
name: Mapped[str] = mapped_column(String(255))
|
||||||
path: Mapped[str] = mapped_column(String(1024))
|
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)
|
owner_id: Mapped[uuid.UUID] = mapped_column(UUID(), ForeignKey("users.id"), nullable=False)
|
||||||
is_mirror: Mapped[bool] = mapped_column(Boolean, default=False, nullable=False)
|
is_mirror: Mapped[bool] = mapped_column(Boolean, default=False, nullable=False)
|
||||||
remote_url: Mapped[str | None] = mapped_column(String(1024), nullable=True)
|
remote_url: Mapped[str | None] = mapped_column(String(1024), nullable=True)
|
||||||
|
|||||||
@@ -0,0 +1,30 @@
|
|||||||
|
"""SQLAlchemy model for health check snapshots."""
|
||||||
|
|
||||||
|
import uuid
|
||||||
|
from datetime import datetime
|
||||||
|
|
||||||
|
from sqlalchemy import Boolean, DateTime, ForeignKey, Integer, String, Text, Uuid, func
|
||||||
|
from sqlalchemy.orm import Mapped, mapped_column
|
||||||
|
|
||||||
|
from src.models.base import Base, UUIDPrimaryKeyMixin
|
||||||
|
|
||||||
|
|
||||||
|
class HealthCheck(UUIDPrimaryKeyMixin, Base):
|
||||||
|
__tablename__ = "health_checks"
|
||||||
|
|
||||||
|
instance_id: Mapped[uuid.UUID] = mapped_column(
|
||||||
|
Uuid(as_uuid=True),
|
||||||
|
ForeignKey("tool_instances.id", ondelete="CASCADE"),
|
||||||
|
nullable=False,
|
||||||
|
)
|
||||||
|
container_status: Mapped[str | None] = mapped_column(String(50), nullable=True)
|
||||||
|
container_healthy: Mapped[bool | None] = mapped_column(Boolean, nullable=True)
|
||||||
|
tunnel_healthy: Mapped[bool | None] = mapped_column(Boolean, nullable=True)
|
||||||
|
exit_code: Mapped[int | None] = mapped_column(Integer, nullable=True)
|
||||||
|
probe_status: Mapped[str | None] = mapped_column(String(50), nullable=True)
|
||||||
|
probe_output: Mapped[str | None] = mapped_column(Text, nullable=True)
|
||||||
|
checked_at: Mapped[datetime] = mapped_column(
|
||||||
|
DateTime(timezone=True),
|
||||||
|
server_default=func.now(),
|
||||||
|
nullable=False,
|
||||||
|
)
|
||||||
@@ -0,0 +1,39 @@
|
|||||||
|
"""SQLAlchemy model for instance lifecycle event audit rows."""
|
||||||
|
|
||||||
|
import uuid
|
||||||
|
from datetime import datetime
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from sqlalchemy import DateTime, ForeignKey, JSON, String, Text, Uuid, func
|
||||||
|
from sqlalchemy.orm import Mapped, mapped_column
|
||||||
|
|
||||||
|
from src.models.base import Base, UUIDPrimaryKeyMixin
|
||||||
|
|
||||||
|
|
||||||
|
class InstanceEvent(UUIDPrimaryKeyMixin, Base):
|
||||||
|
__tablename__ = "instance_events"
|
||||||
|
|
||||||
|
instance_id: Mapped[uuid.UUID] = mapped_column(
|
||||||
|
Uuid(as_uuid=True),
|
||||||
|
ForeignKey("tool_instances.id", ondelete="CASCADE"),
|
||||||
|
nullable=False,
|
||||||
|
)
|
||||||
|
event_type: Mapped[str] = mapped_column(String(50), nullable=False)
|
||||||
|
status: Mapped[str | None] = mapped_column(String(50), nullable=True)
|
||||||
|
message: Mapped[str | None] = mapped_column(Text, nullable=True)
|
||||||
|
created_by: Mapped[uuid.UUID | None] = mapped_column(
|
||||||
|
Uuid(as_uuid=True),
|
||||||
|
ForeignKey("users.id", ondelete="SET NULL"),
|
||||||
|
nullable=True,
|
||||||
|
)
|
||||||
|
event_metadata: Mapped[dict[str, Any]] = mapped_column(
|
||||||
|
"metadata",
|
||||||
|
JSON,
|
||||||
|
nullable=False,
|
||||||
|
default=dict,
|
||||||
|
)
|
||||||
|
created_at: Mapped[datetime] = mapped_column(
|
||||||
|
DateTime(timezone=True),
|
||||||
|
server_default=func.now(),
|
||||||
|
nullable=False,
|
||||||
|
)
|
||||||
@@ -0,0 +1,43 @@
|
|||||||
|
"""Notification SQLAlchemy model."""
|
||||||
|
|
||||||
|
from datetime import datetime
|
||||||
|
from typing import Any
|
||||||
|
import uuid
|
||||||
|
|
||||||
|
from sqlalchemy import DateTime, ForeignKey, JSON, String, Text
|
||||||
|
from sqlalchemy import Uuid as UUID
|
||||||
|
from sqlalchemy.orm import Mapped, mapped_column
|
||||||
|
from sqlalchemy.sql import func
|
||||||
|
|
||||||
|
from src.models.base import Base, UUIDPrimaryKeyMixin
|
||||||
|
|
||||||
|
|
||||||
|
class Notification(UUIDPrimaryKeyMixin, Base):
|
||||||
|
__tablename__ = "notifications"
|
||||||
|
|
||||||
|
user_id: Mapped[uuid.UUID] = mapped_column(
|
||||||
|
UUID(as_uuid=True),
|
||||||
|
ForeignKey("users.id", ondelete="CASCADE"),
|
||||||
|
nullable=False,
|
||||||
|
index=True,
|
||||||
|
)
|
||||||
|
category: Mapped[str] = mapped_column(String(32), nullable=False)
|
||||||
|
severity: Mapped[str] = mapped_column(String(16), nullable=False)
|
||||||
|
title: Mapped[str] = mapped_column(String(255), nullable=False)
|
||||||
|
message: Mapped[str | None] = mapped_column(Text, nullable=True)
|
||||||
|
source_type: Mapped[str | None] = mapped_column(String(64), nullable=True)
|
||||||
|
source_id: Mapped[uuid.UUID | None] = mapped_column(
|
||||||
|
UUID(as_uuid=True), nullable=True
|
||||||
|
)
|
||||||
|
notification_metadata: Mapped[dict[str, Any]] = mapped_column(
|
||||||
|
"metadata", JSON, nullable=False, default=dict
|
||||||
|
)
|
||||||
|
read_at: Mapped[datetime | None] = mapped_column(
|
||||||
|
DateTime(timezone=True), nullable=True, index=True
|
||||||
|
)
|
||||||
|
dismissed_at: Mapped[datetime | None] = mapped_column(
|
||||||
|
DateTime(timezone=True), nullable=True
|
||||||
|
)
|
||||||
|
created_at: Mapped[datetime] = mapped_column(
|
||||||
|
DateTime(timezone=True), server_default=func.now(), nullable=False, index=True
|
||||||
|
)
|
||||||
@@ -0,0 +1,37 @@
|
|||||||
|
"""Terminal session database model."""
|
||||||
|
|
||||||
|
import uuid
|
||||||
|
from datetime import datetime
|
||||||
|
|
||||||
|
from sqlalchemy import DateTime, ForeignKey, String
|
||||||
|
from sqlalchemy import Uuid as UUID
|
||||||
|
from sqlalchemy.orm import Mapped, mapped_column
|
||||||
|
|
||||||
|
from src.models.base import Base, TimestampMixin, UUIDPrimaryKeyMixin
|
||||||
|
|
||||||
|
|
||||||
|
class TerminalSessionModel(UUIDPrimaryKeyMixin, TimestampMixin, Base):
|
||||||
|
"""Database model for terminal session metadata."""
|
||||||
|
|
||||||
|
__tablename__ = "terminal_sessions"
|
||||||
|
|
||||||
|
instance_id: Mapped[uuid.UUID] = mapped_column(
|
||||||
|
UUID(),
|
||||||
|
ForeignKey("tool_instances.id", ondelete="CASCADE"),
|
||||||
|
nullable=False,
|
||||||
|
index=True,
|
||||||
|
)
|
||||||
|
name: Mapped[str | None] = mapped_column(String(255), nullable=True)
|
||||||
|
status: Mapped[str] = mapped_column(
|
||||||
|
String(50),
|
||||||
|
nullable=False,
|
||||||
|
default="active",
|
||||||
|
)
|
||||||
|
last_activity_at: Mapped[datetime | None] = mapped_column(
|
||||||
|
DateTime(timezone=True),
|
||||||
|
nullable=True,
|
||||||
|
)
|
||||||
|
closed_at: Mapped[datetime | None] = mapped_column(
|
||||||
|
DateTime(timezone=True),
|
||||||
|
nullable=True,
|
||||||
|
)
|
||||||
@@ -1,48 +0,0 @@
|
|||||||
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,67 @@
|
|||||||
|
"""Tool Definition Manifest model."""
|
||||||
|
|
||||||
|
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 ToolDefinitionManifest(UUIDPrimaryKeyMixin, TimestampMixin, Base):
|
||||||
|
"""A declarative manifest that compiles to Dockerfile + Compose.
|
||||||
|
|
||||||
|
Can be either:
|
||||||
|
- A base definition (is_base=True) with a FROM image and common packages
|
||||||
|
- A tool definition (is_base=False) that references a base + adds specifics
|
||||||
|
"""
|
||||||
|
|
||||||
|
__tablename__ = "tool_definition_manifests"
|
||||||
|
|
||||||
|
name: Mapped[str] = mapped_column(String(64), unique=True, nullable=False)
|
||||||
|
display_name: Mapped[str] = mapped_column(String(128), nullable=False)
|
||||||
|
description: Mapped[str | None] = mapped_column(Text, nullable=True)
|
||||||
|
category: Mapped[str | None] = mapped_column(String(64), nullable=True)
|
||||||
|
interface_type: Mapped[str] = mapped_column(String(16), nullable=False)
|
||||||
|
|
||||||
|
# Base: either a direct image or a reference to another manifest
|
||||||
|
base_image: Mapped[str | None] = mapped_column(String(256), nullable=True)
|
||||||
|
base_definition_id: Mapped[uuid.UUID | None] = mapped_column(
|
||||||
|
UUID(),
|
||||||
|
ForeignKey("tool_definition_manifests.id"),
|
||||||
|
nullable=True,
|
||||||
|
)
|
||||||
|
base_version: Mapped[str] = mapped_column(
|
||||||
|
String(32), nullable=False, default="latest"
|
||||||
|
)
|
||||||
|
|
||||||
|
# The full manifest JSON
|
||||||
|
manifest: Mapped[dict] = mapped_column(JSON, nullable=False)
|
||||||
|
|
||||||
|
# Caches for quick inspection
|
||||||
|
dockerfile_cache: Mapped[str | None] = mapped_column(Text, nullable=True)
|
||||||
|
compose_cache: Mapped[str | None] = mapped_column(Text, nullable=True)
|
||||||
|
|
||||||
|
# Versioning
|
||||||
|
version: Mapped[str] = mapped_column(String(32), nullable=False, default="v1")
|
||||||
|
is_base: Mapped[bool] = mapped_column(Boolean, nullable=False, default=False)
|
||||||
|
|
||||||
|
created_by_id: Mapped[uuid.UUID | None] = mapped_column(
|
||||||
|
UUID(),
|
||||||
|
ForeignKey("users.id"),
|
||||||
|
nullable=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Relationships
|
||||||
|
created_by: Mapped["User | None"] = relationship(
|
||||||
|
foreign_keys=[created_by_id],
|
||||||
|
)
|
||||||
|
base_definition: Mapped["ToolDefinitionManifest | None"] = relationship(
|
||||||
|
remote_side="ToolDefinitionManifest.id",
|
||||||
|
foreign_keys=[base_definition_id],
|
||||||
|
)
|
||||||
@@ -33,48 +33,33 @@ class ToolInstance(UUIDPrimaryKeyMixin, TimestampMixin, Base):
|
|||||||
owner_id: Mapped[uuid.UUID] = mapped_column(
|
owner_id: Mapped[uuid.UUID] = mapped_column(
|
||||||
UUID(), ForeignKey("users.id"), nullable=False
|
UUID(), ForeignKey("users.id"), nullable=False
|
||||||
)
|
)
|
||||||
status: Mapped[str] = mapped_column(
|
status: Mapped[str] = mapped_column(String(50), nullable=False, default="pending")
|
||||||
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)
|
||||||
container_id: Mapped[str | None] = mapped_column(
|
compose_path: Mapped[str | None] = mapped_column(String(1024), nullable=True)
|
||||||
String(255), nullable=True
|
url: Mapped[str | None] = mapped_column(String(1024), nullable=True)
|
||||||
)
|
public_url: Mapped[str | None] = mapped_column(String(1024), nullable=True)
|
||||||
container_name: Mapped[str | None] = mapped_column(
|
tunnel_id: Mapped[str | None] = mapped_column(String(255), nullable=True)
|
||||||
String(255), nullable=True
|
port: Mapped[int | None] = mapped_column(Integer, 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(
|
last_started_at: Mapped[datetime | None] = mapped_column(
|
||||||
DateTime(timezone=True), nullable=True
|
DateTime(timezone=True), nullable=True
|
||||||
)
|
)
|
||||||
last_stopped_at: Mapped[datetime | None] = mapped_column(
|
last_stopped_at: Mapped[datetime | None] = mapped_column(
|
||||||
DateTime(timezone=True), nullable=True
|
DateTime(timezone=True), nullable=True
|
||||||
)
|
)
|
||||||
probe_result: Mapped[dict | None] = mapped_column(
|
manifest_compiled_at: Mapped[datetime | None] = mapped_column(
|
||||||
JSON, nullable=True
|
DateTime(timezone=True), nullable=True
|
||||||
)
|
|
||||||
clone_mode: Mapped[str] = mapped_column(
|
|
||||||
String(20), nullable=False, default="mount"
|
|
||||||
)
|
)
|
||||||
|
image_tag: Mapped[str | None] = mapped_column(String(256), 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(
|
branch: Mapped[str | None] = mapped_column(
|
||||||
String(255), nullable=True, default="main"
|
String(255), nullable=True, default="main"
|
||||||
)
|
)
|
||||||
selected_config_profile_id: Mapped[uuid.UUID | None] = mapped_column(
|
selected_config_profile_id: Mapped[uuid.UUID | None] = mapped_column(
|
||||||
UUID(), ForeignKey("config_profiles.id", ondelete="SET NULL"), nullable=True
|
UUID(), ForeignKey("config_profiles.id", ondelete="SET NULL"), nullable=True
|
||||||
)
|
)
|
||||||
|
ssh_key_ids: Mapped[list[str] | None] = mapped_column(JSON, nullable=True)
|
||||||
|
|
||||||
tool_type: Mapped["ToolType"] = relationship()
|
tool_type: Mapped["ToolType"] = relationship()
|
||||||
repository: Mapped["GitRepository"] = relationship()
|
repository: Mapped["GitRepository"] = relationship()
|
||||||
|
|||||||
@@ -8,6 +8,7 @@ from sqlalchemy.orm import Mapped, mapped_column, relationship
|
|||||||
from src.models.base import Base, TimestampMixin, UUIDPrimaryKeyMixin
|
from src.models.base import Base, TimestampMixin, UUIDPrimaryKeyMixin
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
|
from src.models.tool_definition_manifest import ToolDefinitionManifest
|
||||||
from src.models.user import User
|
from src.models.user import User
|
||||||
|
|
||||||
|
|
||||||
@@ -18,23 +19,36 @@ class ToolType(UUIDPrimaryKeyMixin, TimestampMixin, Base):
|
|||||||
display_name: Mapped[str] = mapped_column(String(255), nullable=False)
|
display_name: Mapped[str] = mapped_column(String(255), nullable=False)
|
||||||
description: Mapped[str | None] = mapped_column(Text, nullable=True)
|
description: Mapped[str | None] = mapped_column(Text, nullable=True)
|
||||||
category: Mapped[str] = mapped_column(String(50), nullable=False, default="other")
|
category: Mapped[str] = mapped_column(String(50), nullable=False, default="other")
|
||||||
interface_type: Mapped[str] = mapped_column(String(20), nullable=False, default="web")
|
interface_type: Mapped[str] = mapped_column(
|
||||||
|
String(20), nullable=False, default="web"
|
||||||
|
)
|
||||||
requires_port: Mapped[bool] = mapped_column(Boolean, default=True, nullable=False)
|
requires_port: Mapped[bool] = mapped_column(Boolean, default=True, nullable=False)
|
||||||
default_port: Mapped[int] = mapped_column(nullable=False)
|
default_port: Mapped[int] = mapped_column(nullable=False)
|
||||||
definition_type: Mapped[str] = mapped_column(
|
definition_type: Mapped[str] = mapped_column(
|
||||||
String(20), nullable=False, default="compose"
|
String(16), nullable=False, default="legacy"
|
||||||
) # "compose" or "dockerfile"
|
) # "legacy" | "manifest"
|
||||||
|
manifest_id: Mapped[uuid.UUID | None] = mapped_column(
|
||||||
|
UUID(),
|
||||||
|
ForeignKey("tool_definition_manifests.id"),
|
||||||
|
nullable=True,
|
||||||
|
)
|
||||||
compose_template: Mapped[str | None] = mapped_column(Text, nullable=True)
|
compose_template: Mapped[str | None] = mapped_column(Text, nullable=True)
|
||||||
dockerfile_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(
|
build_context: Mapped[dict | None] = mapped_column(
|
||||||
JSON, default=dict, nullable=True
|
JSON, default=dict, nullable=True
|
||||||
)
|
)
|
||||||
readiness_probe: Mapped[dict | None] = mapped_column(JSON, nullable=True)
|
readiness_probe: Mapped[dict | None] = mapped_column(JSON, nullable=True)
|
||||||
required_variables: Mapped[list[str]] = mapped_column(JSON, default=list, nullable=False)
|
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(
|
created_by_id: Mapped[uuid.UUID | None] = mapped_column(
|
||||||
UUID(),
|
UUID(),
|
||||||
ForeignKey("users.id"),
|
ForeignKey("users.id"),
|
||||||
nullable=True,
|
nullable=True,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
manifest: Mapped["ToolDefinitionManifest | None"] = relationship(
|
||||||
|
foreign_keys=[manifest_id],
|
||||||
|
)
|
||||||
created_by: Mapped["User | None"] = relationship()
|
created_by: Mapped["User | None"] = relationship()
|
||||||
|
|||||||
@@ -42,7 +42,7 @@ def clone_repository(
|
|||||||
str(clone_path),
|
str(clone_path),
|
||||||
]
|
]
|
||||||
|
|
||||||
logger.info("Cloning repository %s (branch: %s) into %s", remote_url, branch, clone_path)
|
logger.debug("Cloning repository %s (branch: %s) into %s", remote_url, branch, clone_path)
|
||||||
result = subprocess.run(
|
result = subprocess.run(
|
||||||
cmd,
|
cmd,
|
||||||
capture_output=True,
|
capture_output=True,
|
||||||
@@ -55,7 +55,7 @@ def clone_repository(
|
|||||||
logger.error("Git clone failed: %s", result.stderr)
|
logger.error("Git clone failed: %s", result.stderr)
|
||||||
raise RuntimeError(f"Failed to clone repository: {result.stderr}")
|
raise RuntimeError(f"Failed to clone repository: {result.stderr}")
|
||||||
|
|
||||||
logger.info("Successfully cloned repository into %s", clone_path)
|
logger.debug("Successfully cloned repository into %s", clone_path)
|
||||||
return str(clone_path)
|
return str(clone_path)
|
||||||
|
|
||||||
|
|
||||||
@@ -94,4 +94,4 @@ def remove_clone_directory(instance_dir: str) -> None:
|
|||||||
if clone_path.exists():
|
if clone_path.exists():
|
||||||
import shutil
|
import shutil
|
||||||
shutil.rmtree(clone_path)
|
shutil.rmtree(clone_path)
|
||||||
logger.info("Removed clone directory: %s", clone_path)
|
logger.debug("Removed clone directory: %s", clone_path)
|
||||||
|
|||||||
@@ -5,6 +5,7 @@ and cycle protection.
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
import logging
|
import logging
|
||||||
|
import os
|
||||||
import uuid
|
import uuid
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from typing import Any
|
from typing import Any
|
||||||
@@ -48,6 +49,7 @@ class ResolvedProfile:
|
|||||||
env_vars: dict[str, str] = field(default_factory=dict)
|
env_vars: dict[str, str] = field(default_factory=dict)
|
||||||
runtime_hints: dict[str, Any] = field(default_factory=dict)
|
runtime_hints: dict[str, Any] = field(default_factory=dict)
|
||||||
mounts: dict[str, ResolvedMount] = 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)
|
files: dict[str, str] = field(default_factory=dict)
|
||||||
env_overrides: 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)
|
hint_overrides: dict[str, str] = field(default_factory=dict)
|
||||||
@@ -56,7 +58,9 @@ class ResolvedProfile:
|
|||||||
included_profiles: list[dict[str, Any]] = field(default_factory=list)
|
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:
|
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.
|
"""Detect if adding profile_id to path would create a cycle.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
@@ -168,6 +172,66 @@ def _merge_mounts(
|
|||||||
return result
|
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.
|
||||||
|
|
||||||
|
Entries with the same remote_url + branch have their mappings concatenated.
|
||||||
|
Different repos are kept as separate entries.
|
||||||
|
All entries are normalized to the mappings format.
|
||||||
|
"""
|
||||||
|
result = list(base)
|
||||||
|
# Normalize existing entries to mappings format
|
||||||
|
for i, m in enumerate(result):
|
||||||
|
result[i] = _normalize_git_mount_entry(dict(m))
|
||||||
|
|
||||||
|
# Build lookup by (remote_url, branch)
|
||||||
|
seen = {}
|
||||||
|
for i, m in enumerate(result):
|
||||||
|
key = (m["remote_url"], m.get("branch"))
|
||||||
|
seen[key] = i
|
||||||
|
|
||||||
|
for mount in overlay:
|
||||||
|
mount = _normalize_git_mount_entry(dict(mount))
|
||||||
|
key = (mount["remote_url"], mount.get("branch"))
|
||||||
|
if key in seen:
|
||||||
|
# Same repo+branch: concatenate mappings, dedup by (source_path, target_path)
|
||||||
|
existing = result[seen[key]]
|
||||||
|
existing_sources = {
|
||||||
|
(m["source_path"], m["target_path"])
|
||||||
|
for m in existing.get("mappings", [])
|
||||||
|
}
|
||||||
|
for mapping in mount.get("mappings", []):
|
||||||
|
map_key = (mapping["source_path"], mapping["target_path"])
|
||||||
|
if map_key not in existing_sources:
|
||||||
|
existing["mappings"].append(dict(mapping))
|
||||||
|
existing_sources.add(map_key)
|
||||||
|
else:
|
||||||
|
seen[key] = len(result)
|
||||||
|
result.append(mount)
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
def _normalize_git_mount_entry(entry: dict[str, Any]) -> dict[str, Any]:
|
||||||
|
"""Normalize a git mount entry to the unified mappings format.
|
||||||
|
|
||||||
|
Converts legacy source_path + target_path into a single-entry mappings array.
|
||||||
|
"""
|
||||||
|
entry = dict(entry)
|
||||||
|
if "mappings" not in entry or not entry.get("mappings"):
|
||||||
|
source = entry.get("source_path", ".")
|
||||||
|
target = entry.get("target_path")
|
||||||
|
if target is not None:
|
||||||
|
entry["mappings"] = [{"source_path": source, "target_path": target}]
|
||||||
|
# Remove legacy fields once normalized
|
||||||
|
entry.pop("source_path", None)
|
||||||
|
entry.pop("target_path", None)
|
||||||
|
return entry
|
||||||
|
|
||||||
|
|
||||||
async def _resolve_profile_recursive(
|
async def _resolve_profile_recursive(
|
||||||
session: AsyncSession,
|
session: AsyncSession,
|
||||||
profile_id: uuid.UUID,
|
profile_id: uuid.UUID,
|
||||||
@@ -191,7 +255,9 @@ async def _resolve_profile_recursive(
|
|||||||
"""
|
"""
|
||||||
if _detect_cycle(profile_id, visited, path):
|
if _detect_cycle(profile_id, visited, path):
|
||||||
cycle_path = " -> ".join(str(p) for p in path + [profile_id])
|
cycle_path = " -> ".join(str(p) for p in path + [profile_id])
|
||||||
raise ConfigProfileCycleError(f"Cycle detected in profile includes: {cycle_path}")
|
raise ConfigProfileCycleError(
|
||||||
|
f"Cycle detected in profile includes: {cycle_path}"
|
||||||
|
)
|
||||||
|
|
||||||
profile = await session.get(ConfigProfile, profile_id)
|
profile = await session.get(ConfigProfile, profile_id)
|
||||||
if profile is None:
|
if profile is None:
|
||||||
@@ -218,13 +284,18 @@ async def _resolve_profile_recursive(
|
|||||||
included = await _resolve_profile_recursive(
|
included = await _resolve_profile_recursive(
|
||||||
session, include.included_profile_id, new_visited, new_path
|
session, include.included_profile_id, new_visited, new_path
|
||||||
)
|
)
|
||||||
result.included_profiles.append({
|
result.included_profiles.append(
|
||||||
"id": str(included.profile_id),
|
{
|
||||||
"name": included.profile_name,
|
"id": str(included.profile_id),
|
||||||
})
|
"name": included.profile_name,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
result.env_vars = _merge_env_vars(
|
result.env_vars = _merge_env_vars(
|
||||||
result.env_vars, included.env_vars, result.env_overrides, included.profile_name
|
result.env_vars,
|
||||||
|
included.env_vars,
|
||||||
|
result.env_overrides,
|
||||||
|
included.profile_name,
|
||||||
)
|
)
|
||||||
result.runtime_hints = _merge_runtime_hints(
|
result.runtime_hints = _merge_runtime_hints(
|
||||||
result.runtime_hints,
|
result.runtime_hints,
|
||||||
@@ -244,6 +315,9 @@ async def _resolve_profile_recursive(
|
|||||||
result.mount_overrides,
|
result.mount_overrides,
|
||||||
included.profile_name,
|
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)
|
# Apply the profile's own settings (selected profile overrides includes)
|
||||||
result.env_vars = _merge_env_vars(
|
result.env_vars = _merge_env_vars(
|
||||||
@@ -270,7 +344,11 @@ async def _resolve_profile_recursive(
|
|||||||
result.mount_overrides,
|
result.mount_overrides,
|
||||||
profile.name,
|
profile.name,
|
||||||
)
|
)
|
||||||
|
result.git_mounts = _merge_git_mounts(
|
||||||
|
result.git_mounts,
|
||||||
|
profile.git_mounts or [],
|
||||||
|
profile.name,
|
||||||
|
)
|
||||||
return result
|
return result
|
||||||
|
|
||||||
|
|
||||||
@@ -361,6 +439,7 @@ async def check_include_cycle(
|
|||||||
def apply_resolved_profile(
|
def apply_resolved_profile(
|
||||||
instance_dir: str,
|
instance_dir: str,
|
||||||
resolved: ResolvedProfile,
|
resolved: ResolvedProfile,
|
||||||
|
home_dir: str = "/root",
|
||||||
) -> tuple[dict[str, str], dict[str, str], list[dict], dict[str, Any]]:
|
) -> tuple[dict[str, str], dict[str, str], list[dict], dict[str, Any]]:
|
||||||
"""Apply a resolved profile to an instance directory.
|
"""Apply a resolved profile to an instance directory.
|
||||||
|
|
||||||
@@ -390,14 +469,19 @@ def apply_resolved_profile(
|
|||||||
try:
|
try:
|
||||||
full_path.resolve().relative_to(instance_path.resolve())
|
full_path.resolve().relative_to(instance_path.resolve())
|
||||||
except ValueError:
|
except ValueError:
|
||||||
logger.warning("Profile file path escapes instance directory: %s", file_path)
|
logger.warning(
|
||||||
|
"Profile file path escapes instance directory: %s", file_path
|
||||||
|
)
|
||||||
continue
|
continue
|
||||||
full_path.parent.mkdir(parents=True, exist_ok=True)
|
full_path.parent.mkdir(parents=True, exist_ok=True)
|
||||||
full_path.write_text(content)
|
full_path.write_text(content)
|
||||||
|
|
||||||
# Stage mount files and prepare volume mounts
|
# Stage mount files and prepare volume mounts
|
||||||
for mount in resolved.mounts.values():
|
for mount in resolved.mounts.values():
|
||||||
mount_dir = instance_path / "mounts" / mount.target.lstrip("/").replace("/", "_")
|
expanded_target = expand_container_path(mount.target, home_dir)
|
||||||
|
mount_dir = (
|
||||||
|
instance_path / "mounts" / expanded_target.lstrip("/").replace("/", "_")
|
||||||
|
)
|
||||||
mount_dir.mkdir(parents=True, exist_ok=True)
|
mount_dir.mkdir(parents=True, exist_ok=True)
|
||||||
|
|
||||||
for file_path, content in mount.files.items():
|
for file_path, content in mount.files.items():
|
||||||
@@ -410,15 +494,44 @@ def apply_resolved_profile(
|
|||||||
full_path.parent.mkdir(parents=True, exist_ok=True)
|
full_path.parent.mkdir(parents=True, exist_ok=True)
|
||||||
full_path.write_text(content)
|
full_path.write_text(content)
|
||||||
|
|
||||||
volume_mounts.append({
|
# Mount each file individually so sibling files from other mounts
|
||||||
"source": str(mount_dir),
|
# (e.g. git repo directories) are preserved.
|
||||||
"target": mount.target,
|
file_target = os.path.join(expanded_target, file_path)
|
||||||
"type": "bind",
|
volume_mounts.append(
|
||||||
})
|
{
|
||||||
|
"source": str(full_path),
|
||||||
|
"target": file_target,
|
||||||
|
"type": "bind",
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
||||||
return env_vars, files, volume_mounts, resolved.runtime_hints
|
return env_vars, files, volume_mounts, resolved.runtime_hints
|
||||||
|
|
||||||
|
|
||||||
|
def expand_container_path(path: str, home_dir: str) -> str:
|
||||||
|
"""Expand ~ and $HOME in a container path to the actual home directory.
|
||||||
|
|
||||||
|
Only expands at the start of the path (e.g., ~/foo, $HOME/foo, $HOME).
|
||||||
|
Leaves mid-string occurrences unchanged.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
path: Container path that may contain ~ or $HOME.
|
||||||
|
home_dir: The container's home directory (e.g., /home/user or /root).
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Path with ~ and $HOME expanded.
|
||||||
|
"""
|
||||||
|
if path.startswith("~/"):
|
||||||
|
return os.path.join(home_dir, path[2:])
|
||||||
|
if path == "~":
|
||||||
|
return home_dir
|
||||||
|
if path.startswith("$HOME/"):
|
||||||
|
return home_dir + "/" + path[6:]
|
||||||
|
if path == "$HOME":
|
||||||
|
return home_dir
|
||||||
|
return path
|
||||||
|
|
||||||
|
|
||||||
def resolved_profile_to_dict(resolved: ResolvedProfile) -> dict[str, Any]:
|
def resolved_profile_to_dict(resolved: ResolvedProfile) -> dict[str, Any]:
|
||||||
"""Convert a ResolvedProfile to a plain dict for serialization.
|
"""Convert a ResolvedProfile to a plain dict for serialization.
|
||||||
|
|
||||||
@@ -449,5 +562,6 @@ def resolved_profile_to_dict(resolved: ResolvedProfile) -> dict[str, Any]:
|
|||||||
"files": resolved.file_overrides,
|
"files": resolved.file_overrides,
|
||||||
"mounts": resolved.mount_overrides,
|
"mounts": resolved.mount_overrides,
|
||||||
},
|
},
|
||||||
|
"git_mounts": resolved.git_mounts,
|
||||||
"included_profiles": resolved.included_profiles,
|
"included_profiles": resolved.included_profiles,
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,32 @@
|
|||||||
|
"""Async correlation ID context variable and helpers."""
|
||||||
|
|
||||||
|
import contextvars
|
||||||
|
import uuid
|
||||||
|
|
||||||
|
from fastapi import Request
|
||||||
|
from starlette.middleware.base import BaseHTTPMiddleware
|
||||||
|
|
||||||
|
CORRELATION_ID: contextvars.ContextVar[str] = contextvars.ContextVar("correlation_id")
|
||||||
|
|
||||||
|
|
||||||
|
def get_correlation_id() -> str:
|
||||||
|
"""Return the current correlation ID or generate a new UUID."""
|
||||||
|
try:
|
||||||
|
return CORRELATION_ID.get()
|
||||||
|
except LookupError:
|
||||||
|
return str(uuid.uuid4())
|
||||||
|
|
||||||
|
|
||||||
|
class CorrelationIdMiddleware(BaseHTTPMiddleware):
|
||||||
|
"""Set correlation ID from X-Request-ID header or generate a new UUID."""
|
||||||
|
|
||||||
|
async def dispatch(self, request: Request, call_next):
|
||||||
|
request_id = request.headers.get("X-Request-ID")
|
||||||
|
correlation_id = request_id or str(uuid.uuid4())
|
||||||
|
token = CORRELATION_ID.set(correlation_id)
|
||||||
|
try:
|
||||||
|
response = await call_next(request)
|
||||||
|
response.headers["X-Request-ID"] = correlation_id
|
||||||
|
return response
|
||||||
|
finally:
|
||||||
|
CORRELATION_ID.reset(token)
|
||||||
+243
-36
@@ -1,10 +1,54 @@
|
|||||||
"""Docker service for managing tool instances."""
|
"""Docker service for managing tool instances."""
|
||||||
|
|
||||||
|
import logging
|
||||||
import os
|
import os
|
||||||
|
import re
|
||||||
import subprocess
|
import subprocess
|
||||||
|
import time
|
||||||
|
from collections import Counter
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
def sort_volumes_by_specificity(volumes: list[str]) -> list[str]:
|
||||||
|
"""Sort volume strings so parent paths come before child paths.
|
||||||
|
|
||||||
|
Docker Compose mounts volumes in array order. A later mount at a parent
|
||||||
|
path hides earlier mounts at child paths. By sorting shallow paths first
|
||||||
|
and deep paths last, deeper (more specific) mounts overlay correctly.
|
||||||
|
|
||||||
|
Volume format: source:target or source:target:type
|
||||||
|
|
||||||
|
Args:
|
||||||
|
volumes: List of Docker volume mount strings.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Sorted list with parent paths before child paths.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def _target_depth(vol: str) -> int:
|
||||||
|
parts = vol.split(":")
|
||||||
|
if len(parts) < 2:
|
||||||
|
return 0
|
||||||
|
target = parts[1].rstrip("/")
|
||||||
|
if not target or target == "/":
|
||||||
|
return 0
|
||||||
|
return target.count("/")
|
||||||
|
|
||||||
|
# Detect duplicate targets and warn
|
||||||
|
targets = []
|
||||||
|
for vol in volumes:
|
||||||
|
parts = vol.split(":")
|
||||||
|
targets.append(parts[1] if len(parts) > 1 else "")
|
||||||
|
dupes = [t for t, c in Counter(targets).items() if c > 1]
|
||||||
|
if dupes:
|
||||||
|
logger.warning("Duplicate mount targets detected: %s", dupes)
|
||||||
|
|
||||||
|
# Stable sort: parent paths first, child paths last
|
||||||
|
return sorted(volumes, key=_target_depth)
|
||||||
|
|
||||||
|
|
||||||
def render_compose_template(template: str, variables: dict[str, Any]) -> str:
|
def render_compose_template(template: str, variables: dict[str, Any]) -> str:
|
||||||
"""Render a Docker Compose template with variable substitution.
|
"""Render a Docker Compose template with variable substitution.
|
||||||
@@ -35,6 +79,7 @@ def ensure_instance_directory(instance_id: str, base_path: str | None = None) ->
|
|||||||
"""
|
"""
|
||||||
if base_path is None:
|
if base_path is None:
|
||||||
from src.config import Settings
|
from src.config import Settings
|
||||||
|
|
||||||
base_path = Settings().instance_base_path
|
base_path = Settings().instance_base_path
|
||||||
instance_dir = Path(base_path) / instance_id
|
instance_dir = Path(base_path) / instance_id
|
||||||
instance_dir.mkdir(parents=True, exist_ok=True)
|
instance_dir.mkdir(parents=True, exist_ok=True)
|
||||||
@@ -87,7 +132,7 @@ def write_config_files(instance_dir: str, files: dict[str, str]) -> None:
|
|||||||
full_path.resolve().relative_to(instance_path.resolve())
|
full_path.resolve().relative_to(instance_path.resolve())
|
||||||
except ValueError:
|
except ValueError:
|
||||||
raise ValueError(f"File path '{file_path}' escapes instance directory")
|
raise ValueError(f"File path '{file_path}' escapes instance directory")
|
||||||
|
|
||||||
full_path.parent.mkdir(parents=True, exist_ok=True)
|
full_path.parent.mkdir(parents=True, exist_ok=True)
|
||||||
full_path.write_text(content)
|
full_path.write_text(content)
|
||||||
|
|
||||||
@@ -109,12 +154,12 @@ def execute_compose_command(
|
|||||||
instance_dir = Path(compose_path).parent
|
instance_dir = Path(compose_path).parent
|
||||||
|
|
||||||
cmd = ["docker", "compose", "-f", compose_path]
|
cmd = ["docker", "compose", "-f", compose_path]
|
||||||
|
|
||||||
if env_file:
|
if env_file:
|
||||||
cmd.extend(["--env-file", env_file])
|
cmd.extend(["--env-file", env_file])
|
||||||
|
|
||||||
if action == "up":
|
if action == "up":
|
||||||
cmd.extend(["up", "-d"])
|
cmd.extend(["up", "-d", "--force-recreate"])
|
||||||
elif action == "down":
|
elif action == "down":
|
||||||
cmd.extend(["down", "-v"])
|
cmd.extend(["down", "-v"])
|
||||||
elif action in ("start", "stop", "restart"):
|
elif action in ("start", "stop", "restart"):
|
||||||
@@ -136,14 +181,17 @@ def execute_compose_command(
|
|||||||
def get_container_id(instance_name: str) -> str | None:
|
def get_container_id(instance_name: str) -> str | None:
|
||||||
"""Get the container ID for a compose service.
|
"""Get the container ID for a compose service.
|
||||||
|
|
||||||
|
Searches all containers including stopped/exited ones.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
instance_name: The service name in compose
|
instance_name: The service name in compose
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Container ID or None if not found
|
Container ID or None if not found
|
||||||
"""
|
"""
|
||||||
|
# Docker container names are lowercase internally; normalize to ensure match
|
||||||
result = subprocess.run(
|
result = subprocess.run(
|
||||||
["docker", "ps", "-q", "--filter", f"name={instance_name}"],
|
["docker", "ps", "-a", "-q", "--filter", f"name={instance_name.lower()}"],
|
||||||
capture_output=True,
|
capture_output=True,
|
||||||
text=True,
|
text=True,
|
||||||
)
|
)
|
||||||
@@ -156,14 +204,25 @@ def get_container_id(instance_name: str) -> str | None:
|
|||||||
def get_container_name(instance_name: str) -> str | None:
|
def get_container_name(instance_name: str) -> str | None:
|
||||||
"""Get the full container name for a compose service.
|
"""Get the full container name for a compose service.
|
||||||
|
|
||||||
|
Searches all containers including stopped/exited ones.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
instance_name: The service name in compose
|
instance_name: The service name in compose
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Container name or None if not found
|
Container name or None if not found
|
||||||
"""
|
"""
|
||||||
|
# Docker container names are lowercase internally; normalize to ensure match
|
||||||
result = subprocess.run(
|
result = subprocess.run(
|
||||||
["docker", "ps", "--format", "{{.Names}}", "--filter", f"name={instance_name}"],
|
[
|
||||||
|
"docker",
|
||||||
|
"ps",
|
||||||
|
"-a",
|
||||||
|
"--format",
|
||||||
|
"{{.Names}}",
|
||||||
|
"--filter",
|
||||||
|
f"name={instance_name.lower()}",
|
||||||
|
],
|
||||||
capture_output=True,
|
capture_output=True,
|
||||||
text=True,
|
text=True,
|
||||||
)
|
)
|
||||||
@@ -173,7 +232,9 @@ def get_container_name(instance_name: str) -> str | None:
|
|||||||
return None
|
return None
|
||||||
|
|
||||||
|
|
||||||
def connect_container_to_network(container_name: str, network_name: str = "backend") -> bool:
|
def connect_container_to_network(
|
||||||
|
container_name: str, network_name: str = "backend"
|
||||||
|
) -> bool:
|
||||||
"""Connect a Docker container to an existing network.
|
"""Connect a Docker container to an existing network.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
@@ -198,12 +259,14 @@ def get_container_status(container_id: str) -> dict[str, Any]:
|
|||||||
container_id: Docker container ID
|
container_id: Docker container ID
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Dict with 'status' (running, exited, restarting, not_found),
|
Dict with 'status' (running, exited, restarting, not_found),
|
||||||
'exit_code' (int or None), and 'health' (health status or None)
|
'exit_code' (int or None), and 'health' (health status or None)
|
||||||
"""
|
"""
|
||||||
result = subprocess.run(
|
result = subprocess.run(
|
||||||
[
|
[
|
||||||
"docker", "inspect", "-f",
|
"docker",
|
||||||
|
"inspect",
|
||||||
|
"-f",
|
||||||
"{{.State.Status}}|{{.State.ExitCode}}|{{if .State.Health}}{{.State.Health.Status}}{{else}}none{{end}}",
|
"{{.State.Status}}|{{.State.ExitCode}}|{{if .State.Health}}{{.State.Health.Status}}{{else}}none{{end}}",
|
||||||
container_id,
|
container_id,
|
||||||
],
|
],
|
||||||
@@ -213,12 +276,12 @@ def get_container_status(container_id: str) -> dict[str, Any]:
|
|||||||
|
|
||||||
if result.returncode != 0:
|
if result.returncode != 0:
|
||||||
return {"status": "not_found", "exit_code": None, "health": None}
|
return {"status": "not_found", "exit_code": None, "health": None}
|
||||||
|
|
||||||
parts = result.stdout.strip().split("|")
|
parts = result.stdout.strip().split("|")
|
||||||
status = parts[0] if parts else "unknown"
|
status = parts[0] if parts else "unknown"
|
||||||
exit_code = int(parts[1]) if len(parts) > 1 and parts[1].isdigit() else None
|
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
|
health = parts[2] if len(parts) > 2 and parts[2] != "none" else None
|
||||||
|
|
||||||
return {"status": status, "exit_code": exit_code, "health": health}
|
return {"status": status, "exit_code": exit_code, "health": health}
|
||||||
|
|
||||||
|
|
||||||
@@ -238,13 +301,12 @@ def wait_for_container_running(
|
|||||||
Dict with 'success' (bool), 'status' (str), 'exit_code' (int or None),
|
Dict with 'success' (bool), 'status' (str), 'exit_code' (int or None),
|
||||||
and 'waited_seconds' (float)
|
and 'waited_seconds' (float)
|
||||||
"""
|
"""
|
||||||
import time
|
|
||||||
|
|
||||||
start_time = time.time()
|
start_time = time.time()
|
||||||
|
|
||||||
while time.time() - start_time < timeout:
|
while time.time() - start_time < timeout:
|
||||||
info = get_container_status(container_id)
|
info = get_container_status(container_id)
|
||||||
|
|
||||||
if info["status"] == "running":
|
if info["status"] == "running":
|
||||||
return {
|
return {
|
||||||
"success": True,
|
"success": True,
|
||||||
@@ -252,7 +314,7 @@ def wait_for_container_running(
|
|||||||
"exit_code": None,
|
"exit_code": None,
|
||||||
"waited_seconds": time.time() - start_time,
|
"waited_seconds": time.time() - start_time,
|
||||||
}
|
}
|
||||||
|
|
||||||
if info["status"] == "exited":
|
if info["status"] == "exited":
|
||||||
return {
|
return {
|
||||||
"success": False,
|
"success": False,
|
||||||
@@ -260,7 +322,7 @@ def wait_for_container_running(
|
|||||||
"exit_code": info["exit_code"],
|
"exit_code": info["exit_code"],
|
||||||
"waited_seconds": time.time() - start_time,
|
"waited_seconds": time.time() - start_time,
|
||||||
}
|
}
|
||||||
|
|
||||||
if info["status"] == "not_found":
|
if info["status"] == "not_found":
|
||||||
return {
|
return {
|
||||||
"success": False,
|
"success": False,
|
||||||
@@ -268,9 +330,9 @@ def wait_for_container_running(
|
|||||||
"exit_code": None,
|
"exit_code": None,
|
||||||
"waited_seconds": time.time() - start_time,
|
"waited_seconds": time.time() - start_time,
|
||||||
}
|
}
|
||||||
|
|
||||||
time.sleep(interval)
|
time.sleep(interval)
|
||||||
|
|
||||||
# Timeout reached
|
# Timeout reached
|
||||||
info = get_container_status(container_id)
|
info = get_container_status(container_id)
|
||||||
return {
|
return {
|
||||||
@@ -322,9 +384,83 @@ def find_free_port(start: int = 10000, end: int = 20000) -> int:
|
|||||||
raise RuntimeError(f"No free port found in range {start}-{end}")
|
raise RuntimeError(f"No free port found in range {start}-{end}")
|
||||||
|
|
||||||
|
|
||||||
import subprocess
|
def _check_app_binding(container_name: str, port: int) -> dict[str, str | bool]:
|
||||||
import time
|
"""Diagnose whether the app is bound to 127.0.0.1 or 0.0.0.0.
|
||||||
import re
|
|
||||||
|
Checks from both inside the container (localhost) and outside
|
||||||
|
(via Docker network) to detect binding issues.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Dict with 'internal_ok', 'external_ok', 'internal_status',
|
||||||
|
'external_status', and 'diagnosis'.
|
||||||
|
"""
|
||||||
|
import subprocess
|
||||||
|
|
||||||
|
result: dict[str, Any] = {
|
||||||
|
"internal_ok": False,
|
||||||
|
"external_ok": False,
|
||||||
|
"internal_status": None,
|
||||||
|
"external_status": None,
|
||||||
|
"diagnosis": "unknown",
|
||||||
|
}
|
||||||
|
|
||||||
|
# Check from inside the container (loopback)
|
||||||
|
internal = subprocess.run(
|
||||||
|
[
|
||||||
|
"docker",
|
||||||
|
"exec",
|
||||||
|
container_name,
|
||||||
|
"sh",
|
||||||
|
"-c",
|
||||||
|
f"curl -s -o /dev/null -w '%{{http_code}}' http://localhost:{port}",
|
||||||
|
],
|
||||||
|
capture_output=True,
|
||||||
|
text=True,
|
||||||
|
timeout=5,
|
||||||
|
)
|
||||||
|
if internal.returncode == 0:
|
||||||
|
try:
|
||||||
|
result["internal_status"] = int(internal.stdout.strip())
|
||||||
|
result["internal_ok"] = result["internal_status"] > 0
|
||||||
|
except ValueError:
|
||||||
|
pass
|
||||||
|
|
||||||
|
# Check from outside the container (Docker network)
|
||||||
|
external = subprocess.run(
|
||||||
|
[
|
||||||
|
"curl",
|
||||||
|
"-s",
|
||||||
|
"-o",
|
||||||
|
"/dev/null",
|
||||||
|
"-w",
|
||||||
|
"%{http_code}",
|
||||||
|
f"http://{container_name}:{port}",
|
||||||
|
],
|
||||||
|
capture_output=True,
|
||||||
|
text=True,
|
||||||
|
timeout=5,
|
||||||
|
)
|
||||||
|
if external.returncode == 0:
|
||||||
|
try:
|
||||||
|
result["external_status"] = int(external.stdout.strip())
|
||||||
|
result["external_ok"] = result["external_status"] > 0
|
||||||
|
except ValueError:
|
||||||
|
pass
|
||||||
|
|
||||||
|
# Diagnose binding issue
|
||||||
|
if result["internal_ok"] and not result["external_ok"]:
|
||||||
|
result["diagnosis"] = (
|
||||||
|
f"App appears to be bound to 127.0.0.1:{port} inside the container. "
|
||||||
|
f"It must bind to 0.0.0.0:{port} to be accessible from the tunnel."
|
||||||
|
)
|
||||||
|
elif result["internal_ok"] and result["external_ok"]:
|
||||||
|
result["diagnosis"] = "App is accessible on both interfaces."
|
||||||
|
elif not result["internal_ok"] and not result["external_ok"]:
|
||||||
|
result["diagnosis"] = f"App is not responding on port {port} at all."
|
||||||
|
else:
|
||||||
|
result["diagnosis"] = "Unexpected binding state."
|
||||||
|
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
def start_cloudflared_tunnel(
|
def start_cloudflared_tunnel(
|
||||||
@@ -344,28 +480,77 @@ def start_cloudflared_tunnel(
|
|||||||
Dict with 'url' (the public tunnel URL) and 'pid' (process ID)
|
Dict with 'url' (the public tunnel URL) and 'pid' (process ID)
|
||||||
"""
|
"""
|
||||||
import subprocess
|
import subprocess
|
||||||
import time
|
|
||||||
import re
|
|
||||||
import logging
|
import logging
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
# First verify the container is accessible
|
# First verify the container is accessible from the Docker network
|
||||||
logger.info("Checking connectivity to %s:%d...", container_name, port)
|
logger.info("Checking connectivity to %s:%d...", container_name, port)
|
||||||
for attempt in range(10):
|
accessible = False
|
||||||
|
last_status = None
|
||||||
|
for attempt in range(30): # 30 attempts × 1s = 30s max wait for app startup
|
||||||
check = subprocess.run(
|
check = subprocess.run(
|
||||||
["curl", "-s", "-o", "/dev/null", "-w", "%{http_code}",
|
[
|
||||||
f"http://{container_name}:{port}"],
|
"curl",
|
||||||
|
"-s",
|
||||||
|
"-o",
|
||||||
|
"/dev/null",
|
||||||
|
"-w",
|
||||||
|
"%{http_code}",
|
||||||
|
"--max-time",
|
||||||
|
"3",
|
||||||
|
f"http://{container_name}:{port}",
|
||||||
|
],
|
||||||
capture_output=True,
|
capture_output=True,
|
||||||
text=True,
|
text=True,
|
||||||
timeout=5,
|
timeout=5,
|
||||||
)
|
)
|
||||||
logger.info("Connectivity check %d: http_code=%s", attempt + 1, check.stdout.strip())
|
status_str = check.stdout.strip()
|
||||||
if check.returncode == 0:
|
logger.info(
|
||||||
break
|
"Connectivity check %d/%d: http_code=%s (rc=%d)",
|
||||||
|
attempt + 1,
|
||||||
|
30,
|
||||||
|
status_str,
|
||||||
|
check.returncode,
|
||||||
|
)
|
||||||
|
try:
|
||||||
|
last_status = int(status_str)
|
||||||
|
# Accept 2xx, 3xx, 401, 403 as "app is listening"
|
||||||
|
if last_status in (401, 403) or 200 <= last_status < 400:
|
||||||
|
accessible = True
|
||||||
|
logger.info(
|
||||||
|
"App on %s:%d is ready (HTTP %d)",
|
||||||
|
container_name,
|
||||||
|
port,
|
||||||
|
last_status,
|
||||||
|
)
|
||||||
|
break
|
||||||
|
except ValueError:
|
||||||
|
pass
|
||||||
|
|
||||||
|
if check.returncode != 0:
|
||||||
|
logger.debug(
|
||||||
|
"curl failed: stderr=%s", check.stderr.strip() if check.stderr else ""
|
||||||
|
)
|
||||||
time.sleep(1)
|
time.sleep(1)
|
||||||
else:
|
|
||||||
logger.warning("Container %s:%d not responding to curl checks", container_name, port)
|
if not accessible:
|
||||||
|
logger.warning(
|
||||||
|
"Container %s:%d not responding after 30s (last status: %s). "
|
||||||
|
"Running binding diagnostics...",
|
||||||
|
container_name,
|
||||||
|
port,
|
||||||
|
last_status,
|
||||||
|
)
|
||||||
|
diagnosis = _check_app_binding(container_name, port)
|
||||||
|
logger.warning(
|
||||||
|
"Binding diagnosis: internal=%s (HTTP %s), external=%s (HTTP %s). %s",
|
||||||
|
diagnosis["internal_ok"],
|
||||||
|
diagnosis["internal_status"],
|
||||||
|
diagnosis["external_ok"],
|
||||||
|
diagnosis["external_status"],
|
||||||
|
diagnosis["diagnosis"],
|
||||||
|
)
|
||||||
|
|
||||||
# Run cloudflared in background, capture output
|
# Run cloudflared in background, capture output
|
||||||
logger.info("Starting cloudflared tunnel to http://%s:%d", container_name, port)
|
logger.info("Starting cloudflared tunnel to http://%s:%d", container_name, port)
|
||||||
@@ -381,9 +566,15 @@ def start_cloudflared_tunnel(
|
|||||||
start_time = time.time()
|
start_time = time.time()
|
||||||
url = None
|
url = None
|
||||||
|
|
||||||
|
if proc.stdout is None:
|
||||||
|
proc.terminate()
|
||||||
|
proc.wait(timeout=5)
|
||||||
|
raise RuntimeError("Failed to capture cloudflared output")
|
||||||
|
|
||||||
while time.time() - start_time < timeout:
|
while time.time() - start_time < timeout:
|
||||||
# Read available output
|
# Read available output
|
||||||
import select
|
import select
|
||||||
|
|
||||||
readable, _, _ = select.select([proc.stdout], [], [], 1.0)
|
readable, _, _ = select.select([proc.stdout], [], [], 1.0)
|
||||||
if readable:
|
if readable:
|
||||||
line = proc.stdout.readline()
|
line = proc.stdout.readline()
|
||||||
@@ -410,7 +601,6 @@ def stop_cloudflared_tunnel(pid: str) -> None:
|
|||||||
Args:
|
Args:
|
||||||
pid: Process ID of the cloudflared tunnel
|
pid: Process ID of the cloudflared tunnel
|
||||||
"""
|
"""
|
||||||
import os
|
|
||||||
import signal
|
import signal
|
||||||
|
|
||||||
try:
|
try:
|
||||||
@@ -455,14 +645,23 @@ def check_tunnel_health(url: str, timeout: int = 10) -> dict[str, Any]:
|
|||||||
|
|
||||||
try:
|
try:
|
||||||
result = subprocess.run(
|
result = subprocess.run(
|
||||||
["curl", "-s", "-o", "/dev/null", "-w", "%{http_code}",
|
[
|
||||||
"--max-time", str(timeout), url],
|
"curl",
|
||||||
|
"-s",
|
||||||
|
"-o",
|
||||||
|
"/dev/null",
|
||||||
|
"-w",
|
||||||
|
"%{http_code}",
|
||||||
|
"--max-time",
|
||||||
|
str(timeout),
|
||||||
|
url,
|
||||||
|
],
|
||||||
capture_output=True,
|
capture_output=True,
|
||||||
text=True,
|
text=True,
|
||||||
timeout=timeout + 5,
|
timeout=timeout + 5,
|
||||||
)
|
)
|
||||||
status_code = int(result.stdout.strip())
|
status_code = int(result.stdout.strip())
|
||||||
|
|
||||||
if 200 <= status_code < 400:
|
if 200 <= status_code < 400:
|
||||||
return {
|
return {
|
||||||
"tunnel_status": "healthy",
|
"tunnel_status": "healthy",
|
||||||
@@ -495,7 +694,15 @@ def check_tunnel_health(url: str, timeout: int = 10) -> dict[str, Any]:
|
|||||||
except (ValueError, Exception) as e:
|
except (ValueError, Exception) as e:
|
||||||
error_str = str(e).lower()
|
error_str = str(e).lower()
|
||||||
# Classify connection errors
|
# Classify connection errors
|
||||||
if any(err in error_str for err in ["connection refused", "econnrefused", "could not resolve", "nodename"]):
|
if any(
|
||||||
|
err in error_str
|
||||||
|
for err in [
|
||||||
|
"connection refused",
|
||||||
|
"econnrefused",
|
||||||
|
"could not resolve",
|
||||||
|
"nodename",
|
||||||
|
]
|
||||||
|
):
|
||||||
return {
|
return {
|
||||||
"tunnel_status": "unreachable",
|
"tunnel_status": "unreachable",
|
||||||
"status_code": None,
|
"status_code": None,
|
||||||
|
|||||||
@@ -6,7 +6,9 @@ import subprocess
|
|||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
def build_image(instance_dir: str, dockerfile: str, tag: str, build_context: dict | None = None) -> tuple[int, str, str]:
|
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.
|
"""Build a Docker image from a Dockerfile.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
@@ -18,13 +20,18 @@ def build_image(instance_dir: str, dockerfile: str, tag: str, build_context: dic
|
|||||||
Returns:
|
Returns:
|
||||||
Tuple of (returncode, stdout, stderr)
|
Tuple of (returncode, stdout, stderr)
|
||||||
"""
|
"""
|
||||||
import os
|
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
|
# Defensive: normalise any CRLF that may have crept in from manifest DB
|
||||||
|
# strings — Docker's legacy builder treats \r as a character after the
|
||||||
|
# backslash, breaking RUN continuations and producing
|
||||||
|
# "unknown instruction" errors.
|
||||||
|
dockerfile = dockerfile.replace("\r\n", "\n").replace("\r", "\n")
|
||||||
|
|
||||||
# Write Dockerfile
|
# Write Dockerfile
|
||||||
dockerfile_path = Path(instance_dir) / "Dockerfile"
|
dockerfile_path = Path(instance_dir) / "Dockerfile"
|
||||||
dockerfile_path.write_text(dockerfile)
|
dockerfile_path.write_text(dockerfile, newline="\n")
|
||||||
logger.info("Wrote Dockerfile to %s", dockerfile_path)
|
logger.debug("Wrote Dockerfile to %s (%d bytes)", dockerfile_path, len(dockerfile))
|
||||||
|
|
||||||
# Write build context files
|
# Write build context files
|
||||||
if build_context:
|
if build_context:
|
||||||
@@ -34,19 +41,27 @@ def build_image(instance_dir: str, dockerfile: str, tag: str, build_context: dic
|
|||||||
try:
|
try:
|
||||||
full_path.resolve().relative_to(Path(instance_dir).resolve())
|
full_path.resolve().relative_to(Path(instance_dir).resolve())
|
||||||
except ValueError:
|
except ValueError:
|
||||||
logger.error("Build context file path escapes instance directory: %s", file_path)
|
logger.error(
|
||||||
raise ValueError(f"Build context file path '{file_path}' escapes instance directory")
|
"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.parent.mkdir(parents=True, exist_ok=True)
|
||||||
full_path.write_text(content)
|
normalized = content.replace("\r\n", "\n").replace("\r", "\n")
|
||||||
logger.info("Wrote build context file: %s", full_path)
|
full_path.write_text(normalized, newline="\n")
|
||||||
|
logger.debug("Wrote build context file: %s", full_path)
|
||||||
|
|
||||||
# Build image
|
# Build image
|
||||||
logger.info("Building Docker image with tag: %s", tag)
|
logger.debug("Building Docker image with tag: %s", tag)
|
||||||
cmd = [
|
cmd = [
|
||||||
"docker", "build",
|
"docker",
|
||||||
"-t", tag,
|
"build",
|
||||||
"-f", str(dockerfile_path),
|
"-t",
|
||||||
|
tag,
|
||||||
|
"-f",
|
||||||
|
str(dockerfile_path),
|
||||||
instance_dir,
|
instance_dir,
|
||||||
]
|
]
|
||||||
|
|
||||||
@@ -57,7 +72,7 @@ def build_image(instance_dir: str, dockerfile: str, tag: str, build_context: dic
|
|||||||
text=True,
|
text=True,
|
||||||
timeout=300, # 5 minute timeout for builds
|
timeout=300, # 5 minute timeout for builds
|
||||||
)
|
)
|
||||||
logger.info("Docker build completed: returncode=%d", result.returncode)
|
logger.debug("Docker build completed: returncode=%d", result.returncode)
|
||||||
if result.returncode != 0:
|
if result.returncode != 0:
|
||||||
logger.error("Docker build failed: %s", result.stderr[:1000])
|
logger.error("Docker build failed: %s", result.stderr[:1000])
|
||||||
return result.returncode, result.stdout, result.stderr
|
return result.returncode, result.stdout, result.stderr
|
||||||
|
|||||||
@@ -0,0 +1,97 @@
|
|||||||
|
"""In-memory typed event bus for instance lifecycle and health events."""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import inspect
|
||||||
|
import logging
|
||||||
|
import uuid
|
||||||
|
from collections.abc import Awaitable, Callable
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
InstanceEventPayload = dict[str, Any]
|
||||||
|
EventCallback = Callable[[InstanceEventPayload], Awaitable[None] | None] # noqa: UP044
|
||||||
|
|
||||||
|
|
||||||
|
class InstanceEventBus:
|
||||||
|
"""Singleton in-memory event bus with typed pub/sub and exception isolation."""
|
||||||
|
|
||||||
|
_instance: "InstanceEventBus | None" = None
|
||||||
|
_lock: asyncio.Lock = asyncio.Lock()
|
||||||
|
|
||||||
|
def __init__(self) -> None:
|
||||||
|
self._subscribers: dict[str, list[tuple[str, EventCallback]]] = {}
|
||||||
|
|
||||||
|
def __new__(cls) -> "InstanceEventBus":
|
||||||
|
if cls._instance is None:
|
||||||
|
cls._instance = super().__new__(cls)
|
||||||
|
cls._instance._subscribers = {}
|
||||||
|
return cls._instance
|
||||||
|
|
||||||
|
def _reset_for_testing(self) -> None:
|
||||||
|
"""Clear all subscribers. For test use only."""
|
||||||
|
self._subscribers.clear()
|
||||||
|
|
||||||
|
def subscribe(
|
||||||
|
self,
|
||||||
|
event_type: str,
|
||||||
|
callback: EventCallback,
|
||||||
|
) -> Callable[[], None]:
|
||||||
|
"""Register a callback for an event type.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
event_type: The event type to subscribe to.
|
||||||
|
callback: A sync or async callable that receives the payload.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
An unsubscribe function.
|
||||||
|
"""
|
||||||
|
if event_type not in self._subscribers:
|
||||||
|
self._subscribers[event_type] = []
|
||||||
|
callback_id = str(uuid.uuid4())
|
||||||
|
self._subscribers[event_type].append((callback_id, callback))
|
||||||
|
|
||||||
|
def unsubscribe() -> None:
|
||||||
|
self.unsubscribe(event_type, callback_id)
|
||||||
|
|
||||||
|
return unsubscribe
|
||||||
|
|
||||||
|
def unsubscribe(self, event_type: str, callback_id: str) -> None:
|
||||||
|
"""Remove a specific callback by ID."""
|
||||||
|
if event_type in self._subscribers:
|
||||||
|
self._subscribers[event_type] = [
|
||||||
|
(cid, cb)
|
||||||
|
for cid, cb in self._subscribers[event_type]
|
||||||
|
if cid != callback_id
|
||||||
|
]
|
||||||
|
if not self._subscribers[event_type]:
|
||||||
|
del self._subscribers[event_type]
|
||||||
|
|
||||||
|
def unsubscribe_all(self, event_type: str) -> None:
|
||||||
|
"""Remove all subscribers for an event type."""
|
||||||
|
self._subscribers.pop(event_type, None)
|
||||||
|
|
||||||
|
async def publish(self, event_type: str, payload: InstanceEventPayload) -> None:
|
||||||
|
"""Deliver payload to all subscribers of event_type.
|
||||||
|
|
||||||
|
Also delivers to subscribers registered under the wildcard "*".
|
||||||
|
Exceptions from individual subscribers are caught and logged;
|
||||||
|
delivery continues to remaining subscribers.
|
||||||
|
"""
|
||||||
|
callbacks: list[tuple[str, EventCallback]] = []
|
||||||
|
callbacks.extend(self._subscribers.get(event_type, []))
|
||||||
|
callbacks.extend(self._subscribers.get("*", []))
|
||||||
|
|
||||||
|
for _callback_id, callback in callbacks:
|
||||||
|
try:
|
||||||
|
if inspect.iscoroutinefunction(callback):
|
||||||
|
await callback(payload)
|
||||||
|
else:
|
||||||
|
callback(payload)
|
||||||
|
except Exception:
|
||||||
|
correlation_id = payload.get("correlation_id", "unknown")
|
||||||
|
logger.exception(
|
||||||
|
"Event subscriber failed for %s",
|
||||||
|
event_type,
|
||||||
|
extra={"correlation_id": correlation_id},
|
||||||
|
)
|
||||||
@@ -0,0 +1,253 @@
|
|||||||
|
"""Background health monitor that polls container and tunnel health."""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import logging
|
||||||
|
import uuid
|
||||||
|
from dataclasses import dataclass
|
||||||
|
from datetime import datetime, timezone
|
||||||
|
|
||||||
|
from sqlalchemy import select
|
||||||
|
from sqlalchemy.ext.asyncio import AsyncSession
|
||||||
|
|
||||||
|
from src.database import SessionLocal
|
||||||
|
from src.models.health_check import HealthCheck
|
||||||
|
from src.models.tool_instance import ToolInstance
|
||||||
|
from src.services.correlation import get_correlation_id
|
||||||
|
from src.services.docker import check_tunnel_health, get_container_status
|
||||||
|
from src.services.event_bus import InstanceEventBus, InstanceEventPayload
|
||||||
|
from src.services.notification_service import notification_service
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class HealthSnapshot:
|
||||||
|
"""In-memory snapshot of an instance's health state."""
|
||||||
|
|
||||||
|
container_status: str | None = None
|
||||||
|
container_healthy: bool | None = None
|
||||||
|
tunnel_healthy: bool | None = None
|
||||||
|
exit_code: int | None = None
|
||||||
|
|
||||||
|
|
||||||
|
class HealthMonitor:
|
||||||
|
"""Polls container and tunnel health, publishing events on state changes."""
|
||||||
|
|
||||||
|
POLL_INTERVAL_SECONDS: float = 15.0
|
||||||
|
_MONITORED_STATUSES: set[str] = {"starting", "running", "unhealthy"}
|
||||||
|
|
||||||
|
def __init__(self, event_bus: InstanceEventBus) -> None:
|
||||||
|
self._event_bus = event_bus
|
||||||
|
self._task: asyncio.Task | None = None
|
||||||
|
self._last_known_state: dict[uuid.UUID, HealthSnapshot] = {}
|
||||||
|
|
||||||
|
def start(self) -> None:
|
||||||
|
"""Idempotent start of the background polling task."""
|
||||||
|
if self._task is not None and not self._task.done():
|
||||||
|
return
|
||||||
|
try:
|
||||||
|
loop = asyncio.get_running_loop()
|
||||||
|
self._task = loop.create_task(self._poll_loop())
|
||||||
|
except RuntimeError:
|
||||||
|
pass
|
||||||
|
|
||||||
|
def stop(self) -> None:
|
||||||
|
"""Cancel the background task and clear state."""
|
||||||
|
if self._task is not None and not self._task.done():
|
||||||
|
self._task.cancel()
|
||||||
|
self._last_known_state.clear()
|
||||||
|
self._task = None
|
||||||
|
|
||||||
|
async def _poll_loop(self) -> None:
|
||||||
|
"""Main polling loop."""
|
||||||
|
while True:
|
||||||
|
try:
|
||||||
|
await asyncio.sleep(self.POLL_INTERVAL_SECONDS)
|
||||||
|
await self._run_check_cycle()
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
break
|
||||||
|
except Exception:
|
||||||
|
logger.exception("Health monitor poll loop error")
|
||||||
|
|
||||||
|
async def _run_check_cycle(self) -> None:
|
||||||
|
"""Check all monitored instances in one cycle."""
|
||||||
|
async with SessionLocal() as session:
|
||||||
|
result = await session.execute(
|
||||||
|
select(ToolInstance).where(
|
||||||
|
ToolInstance.status.in_(self._MONITORED_STATUSES)
|
||||||
|
)
|
||||||
|
)
|
||||||
|
instances = result.scalars().all()
|
||||||
|
|
||||||
|
for instance in instances:
|
||||||
|
async with SessionLocal() as session:
|
||||||
|
await self._check_instance(session, instance)
|
||||||
|
|
||||||
|
async def _check_instance(
|
||||||
|
self,
|
||||||
|
session: AsyncSession,
|
||||||
|
instance: ToolInstance,
|
||||||
|
) -> None:
|
||||||
|
"""Check a single instance and handle state transitions."""
|
||||||
|
try:
|
||||||
|
container_info = get_container_status(instance.container_id or "")
|
||||||
|
except Exception:
|
||||||
|
logger.exception(
|
||||||
|
"Health check failed for instance %s",
|
||||||
|
instance.id,
|
||||||
|
extra={
|
||||||
|
"instance_id": str(instance.id),
|
||||||
|
"correlation_id": get_correlation_id(),
|
||||||
|
},
|
||||||
|
)
|
||||||
|
return
|
||||||
|
|
||||||
|
container_status = container_info["status"]
|
||||||
|
exit_code = container_info["exit_code"]
|
||||||
|
container_healthy = (
|
||||||
|
container_info["health"] == "healthy" if container_info["health"] else None
|
||||||
|
)
|
||||||
|
|
||||||
|
tunnel_healthy: bool | None = None
|
||||||
|
if instance.public_url and container_status == "running":
|
||||||
|
try:
|
||||||
|
tunnel_result = check_tunnel_health(instance.public_url)
|
||||||
|
tunnel_healthy = tunnel_result.get("healthy", False)
|
||||||
|
except Exception:
|
||||||
|
logger.exception(
|
||||||
|
"Tunnel health check failed for instance %s",
|
||||||
|
instance.id,
|
||||||
|
extra={
|
||||||
|
"instance_id": str(instance.id),
|
||||||
|
"correlation_id": get_correlation_id(),
|
||||||
|
},
|
||||||
|
)
|
||||||
|
tunnel_healthy = False
|
||||||
|
|
||||||
|
snapshot = HealthSnapshot(
|
||||||
|
container_status=container_status,
|
||||||
|
container_healthy=container_healthy,
|
||||||
|
tunnel_healthy=tunnel_healthy,
|
||||||
|
exit_code=exit_code,
|
||||||
|
)
|
||||||
|
|
||||||
|
previous = self._last_known_state.get(instance.id)
|
||||||
|
|
||||||
|
# Determine new status
|
||||||
|
new_status = self._derive_status(snapshot)
|
||||||
|
|
||||||
|
# If first check or state changed
|
||||||
|
if previous is None or not self._snapshots_equal(previous, snapshot):
|
||||||
|
await self._handle_state_change(
|
||||||
|
session, instance, previous, snapshot, new_status
|
||||||
|
)
|
||||||
|
self._last_known_state[instance.id] = snapshot
|
||||||
|
|
||||||
|
def _derive_status(self, snapshot: HealthSnapshot) -> str:
|
||||||
|
"""Derive instance status from health snapshot."""
|
||||||
|
if snapshot.container_status != "running":
|
||||||
|
return "error"
|
||||||
|
if snapshot.tunnel_healthy is False:
|
||||||
|
return "unhealthy"
|
||||||
|
return "running"
|
||||||
|
|
||||||
|
def _snapshots_equal(self, a: HealthSnapshot, b: HealthSnapshot) -> bool:
|
||||||
|
"""Compare two snapshots for equality."""
|
||||||
|
return (
|
||||||
|
a.container_status == b.container_status
|
||||||
|
and a.container_healthy == b.container_healthy
|
||||||
|
and a.tunnel_healthy == b.tunnel_healthy
|
||||||
|
and a.exit_code == b.exit_code
|
||||||
|
)
|
||||||
|
|
||||||
|
async def _handle_state_change(
|
||||||
|
self,
|
||||||
|
session: AsyncSession,
|
||||||
|
instance: ToolInstance,
|
||||||
|
previous: HealthSnapshot | None,
|
||||||
|
snapshot: HealthSnapshot,
|
||||||
|
new_status: str,
|
||||||
|
) -> None:
|
||||||
|
"""Update DB, insert health check, and publish event."""
|
||||||
|
previous_status = instance.status
|
||||||
|
|
||||||
|
# Update instance status
|
||||||
|
instance.status = new_status
|
||||||
|
if new_status == "error":
|
||||||
|
instance.last_stopped_at = datetime.now(timezone.utc)
|
||||||
|
|
||||||
|
# Insert health check row
|
||||||
|
health_check = HealthCheck(
|
||||||
|
instance_id=instance.id,
|
||||||
|
container_status=snapshot.container_status,
|
||||||
|
container_healthy=snapshot.container_healthy,
|
||||||
|
tunnel_healthy=snapshot.tunnel_healthy,
|
||||||
|
exit_code=snapshot.exit_code,
|
||||||
|
probe_status=None,
|
||||||
|
probe_output=None,
|
||||||
|
)
|
||||||
|
session.add(health_check)
|
||||||
|
await session.commit()
|
||||||
|
|
||||||
|
# Build event payload
|
||||||
|
correlation_id = get_correlation_id()
|
||||||
|
metadata: dict = {"previous_status": previous_status}
|
||||||
|
if snapshot.exit_code is not None:
|
||||||
|
metadata["exit_code"] = snapshot.exit_code
|
||||||
|
metadata["error_type"] = "container"
|
||||||
|
if instance.public_url:
|
||||||
|
metadata["tunnel_url"] = instance.public_url
|
||||||
|
|
||||||
|
if new_status == "error":
|
||||||
|
event_type = "instance.error"
|
||||||
|
message = f"Container failed with status {snapshot.container_status}"
|
||||||
|
if snapshot.exit_code is not None:
|
||||||
|
message += f" (exit code: {snapshot.exit_code})"
|
||||||
|
else:
|
||||||
|
event_type = "instance.health_changed"
|
||||||
|
message = f"Container is now {new_status}"
|
||||||
|
|
||||||
|
payload: InstanceEventPayload = {
|
||||||
|
"event": event_type,
|
||||||
|
"instance_id": str(instance.id),
|
||||||
|
"status": new_status,
|
||||||
|
"message": message,
|
||||||
|
"metadata": metadata,
|
||||||
|
"timestamp": datetime.now(timezone.utc).isoformat(),
|
||||||
|
"correlation_id": correlation_id,
|
||||||
|
}
|
||||||
|
|
||||||
|
await self._event_bus.publish(event_type, payload)
|
||||||
|
|
||||||
|
# Create notification for instance owner (fire-and-forget)
|
||||||
|
# Only send warnings and errors; skip "recovered" info notifications.
|
||||||
|
if new_status == "error":
|
||||||
|
category = "instance"
|
||||||
|
severity = "error"
|
||||||
|
title = "Container failed"
|
||||||
|
elif new_status == "unhealthy":
|
||||||
|
category = "health"
|
||||||
|
severity = "warning"
|
||||||
|
title = "Container unhealthy"
|
||||||
|
else:
|
||||||
|
# Running/recovered — do not notify
|
||||||
|
return
|
||||||
|
|
||||||
|
try:
|
||||||
|
await notification_service.create_notification(
|
||||||
|
session=session,
|
||||||
|
user_id=instance.owner_id,
|
||||||
|
category=category,
|
||||||
|
severity=severity,
|
||||||
|
title=title,
|
||||||
|
message=message,
|
||||||
|
source_type="tool_instances",
|
||||||
|
source_id=instance.id,
|
||||||
|
metadata=metadata,
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
logger.exception(
|
||||||
|
"Failed to create notification for health event %s",
|
||||||
|
event_type,
|
||||||
|
extra={"correlation_id": correlation_id},
|
||||||
|
)
|
||||||
@@ -0,0 +1,162 @@
|
|||||||
|
"""Lifecycle hook helpers for instrumenting tool instance transitions."""
|
||||||
|
|
||||||
|
import logging
|
||||||
|
import uuid
|
||||||
|
from datetime import datetime, timezone
|
||||||
|
|
||||||
|
from sqlalchemy.ext.asyncio import AsyncSession
|
||||||
|
|
||||||
|
from src.models.instance_event import InstanceEvent
|
||||||
|
from src.models.tool_instance import ToolInstance
|
||||||
|
from src.services.correlation import get_correlation_id
|
||||||
|
from src.services.event_bus import InstanceEventBus, InstanceEventPayload
|
||||||
|
from src.services.notification_service import notification_service
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
def _derive_title(event_type: str) -> str:
|
||||||
|
"""Map lifecycle event type to a human-readable notification title."""
|
||||||
|
mapping = {
|
||||||
|
"instance.created": "Container created",
|
||||||
|
"instance.started": "Container started",
|
||||||
|
"instance.stopped": "Container stopped",
|
||||||
|
"instance.restarted": "Container restarted",
|
||||||
|
"instance.deleted": "Container deleted",
|
||||||
|
"instance.error": "Container error",
|
||||||
|
"instance.health_changed": "Container ready",
|
||||||
|
}
|
||||||
|
return mapping.get(
|
||||||
|
event_type,
|
||||||
|
event_type.replace("instance.", "").replace("_", " ").title(),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _should_notify(event_type: str, status: str | None) -> bool:
|
||||||
|
"""Determine whether a lifecycle event should generate a notification.
|
||||||
|
|
||||||
|
Only warnings, errors, and "container is ready" (health_changed running)
|
||||||
|
are sent to users.
|
||||||
|
"""
|
||||||
|
if event_type == "instance.error":
|
||||||
|
return True
|
||||||
|
if event_type == "instance.health_changed" and status == "running":
|
||||||
|
return True
|
||||||
|
# Filter out: created, started, stopped, restarted, deleted, and any
|
||||||
|
# health_changed that is not "running" (unhealthy is handled by health_monitor)
|
||||||
|
return False
|
||||||
|
|
||||||
|
|
||||||
|
def _build_payload(
|
||||||
|
event_type: str,
|
||||||
|
instance: ToolInstance,
|
||||||
|
status: str | None = None,
|
||||||
|
message: str | None = None,
|
||||||
|
metadata: dict | None = None,
|
||||||
|
) -> InstanceEventPayload:
|
||||||
|
"""Construct a standard event payload."""
|
||||||
|
return {
|
||||||
|
"event": event_type,
|
||||||
|
"instance_id": str(instance.id),
|
||||||
|
"status": status or instance.status,
|
||||||
|
"message": message,
|
||||||
|
"metadata": metadata or {},
|
||||||
|
"timestamp": datetime.now(timezone.utc).isoformat(),
|
||||||
|
"correlation_id": get_correlation_id(),
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
async def _write_audit_row(
|
||||||
|
session: AsyncSession,
|
||||||
|
instance: ToolInstance,
|
||||||
|
event_type: str,
|
||||||
|
created_by: uuid.UUID | None = None,
|
||||||
|
status: str | None = None,
|
||||||
|
message: str | None = None,
|
||||||
|
metadata: dict | None = None,
|
||||||
|
) -> InstanceEvent:
|
||||||
|
"""Persist an instance_events audit row."""
|
||||||
|
row = InstanceEvent(
|
||||||
|
instance_id=instance.id,
|
||||||
|
event_type=event_type.replace("instance.", ""),
|
||||||
|
status=status or instance.status,
|
||||||
|
message=message,
|
||||||
|
created_by=created_by,
|
||||||
|
event_metadata=metadata or {},
|
||||||
|
)
|
||||||
|
session.add(row)
|
||||||
|
await session.commit()
|
||||||
|
return row
|
||||||
|
|
||||||
|
|
||||||
|
async def publish_lifecycle_event(
|
||||||
|
event_bus: InstanceEventBus,
|
||||||
|
session: AsyncSession,
|
||||||
|
instance: ToolInstance,
|
||||||
|
event_type: str,
|
||||||
|
created_by: uuid.UUID | None = None,
|
||||||
|
status: str | None = None,
|
||||||
|
message: str | None = None,
|
||||||
|
metadata: dict | None = None,
|
||||||
|
) -> None:
|
||||||
|
"""Publish a lifecycle event and write an audit row after DB commit.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
event_bus: The global event bus.
|
||||||
|
session: Active async DB session.
|
||||||
|
instance: The affected tool instance.
|
||||||
|
event_type: One of instance.created, instance.started, etc.
|
||||||
|
created_by: User ID for user-initiated actions; None for system.
|
||||||
|
status: Optional status override.
|
||||||
|
message: Optional human-readable message.
|
||||||
|
metadata: Optional extra metadata.
|
||||||
|
"""
|
||||||
|
payload = _build_payload(
|
||||||
|
event_type=event_type,
|
||||||
|
instance=instance,
|
||||||
|
status=status,
|
||||||
|
message=message,
|
||||||
|
metadata=metadata,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Write audit row
|
||||||
|
await _write_audit_row(
|
||||||
|
session=session,
|
||||||
|
instance=instance,
|
||||||
|
event_type=event_type,
|
||||||
|
created_by=created_by,
|
||||||
|
status=status or instance.status,
|
||||||
|
message=message,
|
||||||
|
metadata=metadata,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Publish to bus
|
||||||
|
await event_bus.publish(event_type, payload)
|
||||||
|
|
||||||
|
# Create notification for instance owner (fire-and-forget)
|
||||||
|
# Only send warnings, errors, and "container is ready" notifications.
|
||||||
|
effective_status = status or instance.status
|
||||||
|
if not _should_notify(event_type, effective_status):
|
||||||
|
return
|
||||||
|
|
||||||
|
severity = "error" if event_type == "instance.error" else "success"
|
||||||
|
title = _derive_title(event_type)
|
||||||
|
|
||||||
|
try:
|
||||||
|
await notification_service.create_notification(
|
||||||
|
session=session,
|
||||||
|
user_id=instance.owner_id,
|
||||||
|
category="instance",
|
||||||
|
severity=severity,
|
||||||
|
title=title,
|
||||||
|
message=message,
|
||||||
|
source_type="tool_instances",
|
||||||
|
source_id=instance.id,
|
||||||
|
metadata=metadata,
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
logger.exception(
|
||||||
|
"Failed to create notification for lifecycle event %s",
|
||||||
|
event_type,
|
||||||
|
extra={"correlation_id": payload.get("correlation_id", "unknown")},
|
||||||
|
)
|
||||||
@@ -0,0 +1,411 @@
|
|||||||
|
"""Manifest compiler: transforms ToolDefinitionManifest into Dockerfile + Compose."""
|
||||||
|
|
||||||
|
import hashlib
|
||||||
|
import json
|
||||||
|
import shlex
|
||||||
|
from copy import deepcopy
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
import yaml
|
||||||
|
|
||||||
|
from src.services.docker import sort_volumes_by_specificity
|
||||||
|
|
||||||
|
|
||||||
|
def resolve_base(manifest: dict) -> dict:
|
||||||
|
"""Merge a base definition into a tool manifest.
|
||||||
|
|
||||||
|
If the manifest has base_definition_id, the base manifest is loaded
|
||||||
|
and merged. Tool-specific values override base values.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
manifest: The tool manifest JSON (may reference a base)
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
A fully resolved manifest with base values merged in.
|
||||||
|
"""
|
||||||
|
result = deepcopy(manifest)
|
||||||
|
|
||||||
|
base_definition_id = result.pop("base_definition_id", None)
|
||||||
|
base_version = result.pop("base_version", "latest")
|
||||||
|
|
||||||
|
if base_definition_id:
|
||||||
|
# This will be provided by the caller (they have the DB session)
|
||||||
|
# For now, we assume the manifest has been pre-resolved
|
||||||
|
# or the caller provides the base manifest separately.
|
||||||
|
pass
|
||||||
|
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
def deep_merge(base: dict, override: dict) -> dict:
|
||||||
|
"""Deep merge two manifests. Arrays are concatenated; dicts are merged.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
base: The base manifest.
|
||||||
|
override: The tool-specific overrides.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Merged manifest.
|
||||||
|
"""
|
||||||
|
merged = deepcopy(base)
|
||||||
|
|
||||||
|
for key, value in override.items():
|
||||||
|
if key == "mounts" and isinstance(value, list):
|
||||||
|
# Concatenate mount arrays
|
||||||
|
existing = merged.get("mounts", [])
|
||||||
|
merged["mounts"] = existing + deepcopy(value)
|
||||||
|
elif key == "scripts" and isinstance(value, dict):
|
||||||
|
# Merge script categories
|
||||||
|
if "scripts" not in merged:
|
||||||
|
merged["scripts"] = {}
|
||||||
|
for script_key, script_value in value.items():
|
||||||
|
existing = merged["scripts"].get(script_key, [])
|
||||||
|
merged["scripts"][script_key] = existing + deepcopy(script_value)
|
||||||
|
elif key == "packages" and isinstance(value, dict):
|
||||||
|
# Union package arrays
|
||||||
|
if "packages" not in merged:
|
||||||
|
merged["packages"] = {}
|
||||||
|
for pkg_key, pkg_value in value.items():
|
||||||
|
if (
|
||||||
|
pkg_key in merged["packages"]
|
||||||
|
and isinstance(merged["packages"][pkg_key], list)
|
||||||
|
and isinstance(pkg_value, list)
|
||||||
|
):
|
||||||
|
merged["packages"][pkg_key] = merged["packages"][
|
||||||
|
pkg_key
|
||||||
|
] + deepcopy(pkg_value)
|
||||||
|
else:
|
||||||
|
merged["packages"][pkg_key] = deepcopy(pkg_value)
|
||||||
|
elif key == "env" and isinstance(value, dict):
|
||||||
|
# Dict merge: override wins on key conflict
|
||||||
|
if "env" not in merged:
|
||||||
|
merged["env"] = {}
|
||||||
|
merged["env"].update(deepcopy(value))
|
||||||
|
elif (
|
||||||
|
isinstance(value, dict) and key in merged and isinstance(merged[key], dict)
|
||||||
|
):
|
||||||
|
# Generic dict merge
|
||||||
|
merged[key] = {**merged[key], **deepcopy(value)}
|
||||||
|
else:
|
||||||
|
# Override entirely
|
||||||
|
merged[key] = deepcopy(value)
|
||||||
|
|
||||||
|
return merged
|
||||||
|
|
||||||
|
|
||||||
|
def compile_dockerfile(manifest: dict) -> str:
|
||||||
|
"""Compile a resolved manifest into a Dockerfile string.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
manifest: Fully resolved manifest JSON.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Dockerfile content.
|
||||||
|
"""
|
||||||
|
lines: list[str] = []
|
||||||
|
|
||||||
|
# FROM
|
||||||
|
base_image = manifest.get("base_image", "ubuntu:24.04")
|
||||||
|
lines.append(f"FROM {base_image}")
|
||||||
|
lines.append("")
|
||||||
|
|
||||||
|
# Build-time environment
|
||||||
|
env = manifest.get("env", {})
|
||||||
|
for key, value in env.items():
|
||||||
|
lines.append(f"ENV {key}={shlex.quote(value)}")
|
||||||
|
if env:
|
||||||
|
lines.append("")
|
||||||
|
|
||||||
|
# System packages (apt)
|
||||||
|
apt_packages = manifest.get("packages", {}).get("apt", [])
|
||||||
|
if apt_packages:
|
||||||
|
lines.append("RUN apt-get update && apt-get install -y \\")
|
||||||
|
for pkg in apt_packages[:-1]:
|
||||||
|
lines.append(f" {pkg} \\")
|
||||||
|
lines.append(f" {apt_packages[-1]} \\")
|
||||||
|
lines.append(" && rm -rf /var/lib/apt/lists/*")
|
||||||
|
lines.append("")
|
||||||
|
|
||||||
|
# Node.js
|
||||||
|
node = manifest.get("packages", {}).get("node")
|
||||||
|
if node:
|
||||||
|
version = node.get("version", "20")
|
||||||
|
lines.append(
|
||||||
|
f"RUN curl -fsSL https://deb.nodesource.com/setup_{version}.x | bash - && \\"
|
||||||
|
)
|
||||||
|
lines.append(" apt-get install -y nodejs && \\")
|
||||||
|
lines.append(" rm -rf /var/lib/apt/lists/*")
|
||||||
|
lines.append("")
|
||||||
|
|
||||||
|
# NPM global packages
|
||||||
|
npm_packages = manifest.get("packages", {}).get("npm_global", [])
|
||||||
|
if npm_packages:
|
||||||
|
pkg_list = " ".join(shlex.quote(p) for p in npm_packages)
|
||||||
|
lines.append(f"RUN npm install -g {pkg_list}")
|
||||||
|
lines.append("")
|
||||||
|
|
||||||
|
# Pip packages
|
||||||
|
pip_packages = manifest.get("packages", {}).get("pip", [])
|
||||||
|
if pip_packages:
|
||||||
|
pkg_list = " ".join(shlex.quote(p) for p in pip_packages)
|
||||||
|
lines.append(f"RUN pip install {pkg_list}")
|
||||||
|
lines.append("")
|
||||||
|
|
||||||
|
# User creation
|
||||||
|
user = manifest.get("user")
|
||||||
|
if user:
|
||||||
|
name = user["name"]
|
||||||
|
uid = user["uid"]
|
||||||
|
gid = user["gid"]
|
||||||
|
create_home = "-m " if user.get("create_home", True) else ""
|
||||||
|
shell = user.get("shell", "/bin/bash")
|
||||||
|
lines.append(f"RUN groupadd -g {gid} {name} && \\")
|
||||||
|
lines.append(f" useradd -u {uid} -g {gid} {create_home}-s {shell} {name}")
|
||||||
|
lines.append("")
|
||||||
|
# Set HOME and USER for runtime compatibility
|
||||||
|
home = f"/home/{name}"
|
||||||
|
lines.append(f"ENV HOME={home}")
|
||||||
|
lines.append(f"ENV USER={name}")
|
||||||
|
lines.append("")
|
||||||
|
|
||||||
|
# Build scripts
|
||||||
|
build_scripts = manifest.get("scripts", {}).get("build", [])
|
||||||
|
for script in build_scripts:
|
||||||
|
# Normalize multi-line scripts into single RUN command
|
||||||
|
stripped_lines = [
|
||||||
|
line.strip() for line in script.strip().split("\n") if line.strip()
|
||||||
|
]
|
||||||
|
if stripped_lines:
|
||||||
|
normalized = " && ".join(stripped_lines)
|
||||||
|
lines.append(f"RUN {normalized}")
|
||||||
|
if build_scripts:
|
||||||
|
lines.append("")
|
||||||
|
|
||||||
|
# Create mount target directories
|
||||||
|
mounts = manifest.get("mounts", [])
|
||||||
|
if mounts:
|
||||||
|
dirs = [mount["target"] for mount in mounts]
|
||||||
|
dir_str = " ".join(dirs)
|
||||||
|
lines.append(f"RUN mkdir -p {dir_str}")
|
||||||
|
if user:
|
||||||
|
lines.append(f"RUN chown -R {user['name']}:{user['name']} {dir_str}")
|
||||||
|
lines.append("")
|
||||||
|
|
||||||
|
# Entrypoint for startup scripts
|
||||||
|
startup_scripts = manifest.get("scripts", {}).get("startup", [])
|
||||||
|
if startup_scripts:
|
||||||
|
lines.append(
|
||||||
|
"COPY .headquarter/entrypoint.sh /usr/local/bin/headquarter-entrypoint"
|
||||||
|
)
|
||||||
|
lines.append("RUN chmod +x /usr/local/bin/headquarter-entrypoint")
|
||||||
|
lines.append("")
|
||||||
|
|
||||||
|
# Switch to runtime user
|
||||||
|
if user:
|
||||||
|
lines.append(f"USER {user['name']}")
|
||||||
|
lines.append(f"WORKDIR /home/{user['name']}")
|
||||||
|
lines.append("")
|
||||||
|
|
||||||
|
# Entrypoint and CMD
|
||||||
|
runtime = manifest.get("runtime", {})
|
||||||
|
if startup_scripts:
|
||||||
|
lines.append('ENTRYPOINT ["/usr/local/bin/headquarter-entrypoint"]')
|
||||||
|
|
||||||
|
command = runtime.get("command", ["/bin/bash"])
|
||||||
|
cmd_json = json.dumps(command)
|
||||||
|
lines.append(f"CMD {cmd_json}")
|
||||||
|
|
||||||
|
return "\n".join(lines)
|
||||||
|
|
||||||
|
|
||||||
|
def compile_entrypoint(manifest: dict) -> str:
|
||||||
|
"""Generate the startup entrypoint script from startup scripts.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
manifest: Fully resolved manifest JSON.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Shell script content.
|
||||||
|
"""
|
||||||
|
lines = ["#!/bin/bash", "set -e", ""]
|
||||||
|
|
||||||
|
startup_scripts = manifest.get("scripts", {}).get("startup", [])
|
||||||
|
for script in startup_scripts:
|
||||||
|
lines.append(script)
|
||||||
|
lines.append("")
|
||||||
|
|
||||||
|
lines.append('exec "$@"')
|
||||||
|
return "\n".join(lines)
|
||||||
|
|
||||||
|
|
||||||
|
def compile_compose(manifest: dict, variables: dict[str, Any]) -> str:
|
||||||
|
"""Compile a resolved manifest into a Docker Compose string.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
manifest: Fully resolved manifest JSON.
|
||||||
|
variables: Resolved values: IMAGE_TAG, INSTANCE_NAME, REPO_PATH, etc.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Docker Compose YAML content.
|
||||||
|
"""
|
||||||
|
runtime = manifest.get("runtime", {})
|
||||||
|
user = manifest.get("user")
|
||||||
|
interface_type = manifest["interface_type"]
|
||||||
|
|
||||||
|
service: dict[str, Any] = {
|
||||||
|
"image": variables["IMAGE_TAG"],
|
||||||
|
"container_name": variables["INSTANCE_NAME"],
|
||||||
|
"restart": "unless-stopped",
|
||||||
|
}
|
||||||
|
|
||||||
|
# Terminal-specific fields
|
||||||
|
if runtime.get("stdin_open", False):
|
||||||
|
service["stdin_open"] = True
|
||||||
|
if runtime.get("tty", False):
|
||||||
|
service["tty"] = True
|
||||||
|
if runtime.get("working_dir"):
|
||||||
|
service["working_dir"] = runtime["working_dir"]
|
||||||
|
|
||||||
|
# User override
|
||||||
|
if user:
|
||||||
|
service["user"] = f"{user['uid']}:{user['gid']}"
|
||||||
|
|
||||||
|
# Ports for web tools
|
||||||
|
default_port = manifest.get("default_port")
|
||||||
|
if interface_type == "web" and default_port:
|
||||||
|
service["ports"] = [f"{variables['TOOL_PORT']}:{default_port}"]
|
||||||
|
|
||||||
|
# Environment
|
||||||
|
env = manifest.get("env", {})
|
||||||
|
if env:
|
||||||
|
service["environment"] = dict(env)
|
||||||
|
|
||||||
|
# Merge extra env from config
|
||||||
|
extra_env = variables.get("EXTRA_ENV", {})
|
||||||
|
if extra_env:
|
||||||
|
if "environment" not in service:
|
||||||
|
service["environment"] = {}
|
||||||
|
service["environment"].update(extra_env)
|
||||||
|
|
||||||
|
# Volumes from mount schema
|
||||||
|
volumes = []
|
||||||
|
for mount in manifest.get("mounts", []):
|
||||||
|
source = resolve_mount_source(mount, variables)
|
||||||
|
if not source:
|
||||||
|
continue
|
||||||
|
target = mount["target"]
|
||||||
|
readonly = ":ro" if mount.get("readonly", False) else ""
|
||||||
|
volumes.append(f"{source}:{target}{readonly}")
|
||||||
|
|
||||||
|
# Append extra volumes from tool config / config profile
|
||||||
|
for vol in variables.get("EXTRA_VOLUMES", []):
|
||||||
|
vol_str = f"{vol['source']}:{vol['target']}"
|
||||||
|
if vol.get("readonly"):
|
||||||
|
vol_str += ":ro"
|
||||||
|
volumes.append(vol_str)
|
||||||
|
|
||||||
|
if volumes:
|
||||||
|
service["volumes"] = sort_volumes_by_specificity(volumes)
|
||||||
|
|
||||||
|
compose = {"services": {"app": service}}
|
||||||
|
return yaml.dump(compose, default_flow_style=False)
|
||||||
|
|
||||||
|
|
||||||
|
def resolve_mount_source(mount: dict, variables: dict[str, Any]) -> str:
|
||||||
|
"""Resolve a mount's source_type to an actual host path.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
mount: Mount definition from manifest.
|
||||||
|
variables: Resolved variables dict.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Host path string, or empty string if unresolved.
|
||||||
|
"""
|
||||||
|
source_type = mount.get("source_type", "host_path")
|
||||||
|
|
||||||
|
if source_type == "repo":
|
||||||
|
return variables.get("REPO_PATH", "")
|
||||||
|
elif source_type == "ssh_key":
|
||||||
|
return variables.get("SSH_PATH", "")
|
||||||
|
elif source_type == "instance":
|
||||||
|
instance_dir = variables.get("INSTANCE_DIR", "")
|
||||||
|
mount_name = mount.get("name", "unknown")
|
||||||
|
return f"{instance_dir}/mounts/{mount_name}"
|
||||||
|
elif source_type == "git_mount":
|
||||||
|
ref = mount.get("git_mount_ref", "default")
|
||||||
|
return variables.get(f"GIT_MOUNT_{ref}", "")
|
||||||
|
elif source_type == "host_path":
|
||||||
|
return mount.get("source", "")
|
||||||
|
|
||||||
|
return ""
|
||||||
|
|
||||||
|
|
||||||
|
def get_manifest_home_dir(manifest: dict) -> str:
|
||||||
|
"""Get the home directory for a container based on manifest user config.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
manifest: Fully resolved manifest JSON.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Home directory path (e.g., /home/user or /root).
|
||||||
|
"""
|
||||||
|
user = manifest.get("user")
|
||||||
|
if user and user.get("name"):
|
||||||
|
return f"/home/{user['name']}"
|
||||||
|
return "/root"
|
||||||
|
|
||||||
|
|
||||||
|
def compute_image_tag(tool_name: str, manifest: dict) -> str:
|
||||||
|
"""Compute a deterministic image tag from manifest content.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
tool_name: Human-readable tool name.
|
||||||
|
manifest: Fully resolved manifest JSON.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Docker image tag string.
|
||||||
|
"""
|
||||||
|
# Canonicalize: sort keys, stable JSON
|
||||||
|
canonical = json.dumps(manifest, sort_keys=True, separators=(",", ":"))
|
||||||
|
hash_suffix = hashlib.sha256(canonical.encode()).hexdigest()[:8]
|
||||||
|
safe_name = tool_name.lower().replace(" ", "-").replace("_", "-")
|
||||||
|
return f"headquarter/{safe_name}-{hash_suffix}:latest"
|
||||||
|
|
||||||
|
|
||||||
|
def merge_with_config(manifest: dict, profile: dict | None = None) -> dict:
|
||||||
|
"""Merge ConfigProfile overrides into a manifest.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
manifest: Base manifest from tool definition.
|
||||||
|
profile: Resolved ConfigProfile (optional).
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Manifest with overrides applied.
|
||||||
|
"""
|
||||||
|
result = deepcopy(manifest)
|
||||||
|
|
||||||
|
extra_env: dict[str, str] = {}
|
||||||
|
extra_volumes: list[dict] = []
|
||||||
|
|
||||||
|
# Apply ConfigProfile
|
||||||
|
if profile:
|
||||||
|
if profile.get("environment_variables"):
|
||||||
|
extra_env.update(profile["environment_variables"])
|
||||||
|
if profile.get("mounts"):
|
||||||
|
extra_volumes.extend(profile["mounts"])
|
||||||
|
# Profile hints override everything
|
||||||
|
hints = profile.get("hints", {})
|
||||||
|
if hints.get("start_command"):
|
||||||
|
result["runtime"] = result.get("runtime", {})
|
||||||
|
result["runtime"]["command"] = hints["start_command"].split()
|
||||||
|
if hints.get("working_directory"):
|
||||||
|
result["runtime"] = result.get("runtime", {})
|
||||||
|
result["runtime"]["working_dir"] = hints["working_directory"]
|
||||||
|
if hints.get("port_override"):
|
||||||
|
result["default_port"] = hints["port_override"]
|
||||||
|
|
||||||
|
# Store merged extras for the compose compiler
|
||||||
|
result["_extra_env"] = extra_env
|
||||||
|
result["_extra_volumes"] = extra_volumes
|
||||||
|
|
||||||
|
return result
|
||||||
@@ -0,0 +1,272 @@
|
|||||||
|
"""Notification persistence service."""
|
||||||
|
|
||||||
|
import uuid
|
||||||
|
from datetime import datetime, timezone
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from sqlalchemy import func, select, update
|
||||||
|
from sqlalchemy.engine import CursorResult
|
||||||
|
from sqlalchemy.ext.asyncio import AsyncSession
|
||||||
|
|
||||||
|
from src.models.notification import Notification
|
||||||
|
|
||||||
|
|
||||||
|
class NotificationService:
|
||||||
|
"""Singleton notification persistence service.
|
||||||
|
|
||||||
|
All methods filter by user_id to enforce strict ownership isolation.
|
||||||
|
"""
|
||||||
|
|
||||||
|
async def create_notification(
|
||||||
|
self,
|
||||||
|
session: AsyncSession,
|
||||||
|
user_id: uuid.UUID,
|
||||||
|
*,
|
||||||
|
category: str,
|
||||||
|
severity: str,
|
||||||
|
title: str,
|
||||||
|
message: str | None = None,
|
||||||
|
source_type: str | None = None,
|
||||||
|
source_id: uuid.UUID | None = None,
|
||||||
|
metadata: dict[str, Any] | None = None,
|
||||||
|
) -> Notification:
|
||||||
|
"""Insert a new notification row.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
session: Database session.
|
||||||
|
user_id: Owner of the notification.
|
||||||
|
category: Notification category (e.g., instance, system, health).
|
||||||
|
severity: Severity level (e.g., info, warning, error, success).
|
||||||
|
title: Short notification title.
|
||||||
|
message: Optional longer message body.
|
||||||
|
source_type: Optional source entity type.
|
||||||
|
source_id: Optional source entity UUID.
|
||||||
|
metadata: Optional JSON metadata dictionary.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
The newly created Notification instance.
|
||||||
|
"""
|
||||||
|
notification = Notification(
|
||||||
|
user_id=user_id,
|
||||||
|
category=category,
|
||||||
|
severity=severity,
|
||||||
|
title=title,
|
||||||
|
message=message,
|
||||||
|
source_type=source_type,
|
||||||
|
source_id=source_id,
|
||||||
|
notification_metadata=metadata or {},
|
||||||
|
)
|
||||||
|
session.add(notification)
|
||||||
|
await session.commit()
|
||||||
|
await session.refresh(notification)
|
||||||
|
return notification
|
||||||
|
|
||||||
|
async def list_notifications(
|
||||||
|
self,
|
||||||
|
session: AsyncSession,
|
||||||
|
user_id: uuid.UUID,
|
||||||
|
*,
|
||||||
|
limit: int = 20,
|
||||||
|
offset: int = 0,
|
||||||
|
unread_only: bool = False,
|
||||||
|
mute_categories: list[str] | None = None,
|
||||||
|
) -> tuple[list[Notification], int]:
|
||||||
|
"""Return paginated notifications for a user.
|
||||||
|
|
||||||
|
Excludes dismissed notifications and applies optional filtering.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
session: Database session.
|
||||||
|
user_id: Owner of the notifications.
|
||||||
|
limit: Maximum number of items to return.
|
||||||
|
offset: Number of items to skip.
|
||||||
|
unread_only: If True, only return unread notifications.
|
||||||
|
mute_categories: Categories to exclude from results.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
A tuple of (items, total_count).
|
||||||
|
"""
|
||||||
|
where_clauses = [
|
||||||
|
Notification.user_id == user_id,
|
||||||
|
Notification.dismissed_at.is_(None),
|
||||||
|
]
|
||||||
|
|
||||||
|
if unread_only:
|
||||||
|
where_clauses.append(Notification.read_at.is_(None))
|
||||||
|
|
||||||
|
if mute_categories:
|
||||||
|
where_clauses.append(Notification.category.not_in(mute_categories))
|
||||||
|
|
||||||
|
total_stmt = (
|
||||||
|
select(func.count()).select_from(Notification).where(*where_clauses)
|
||||||
|
)
|
||||||
|
total_result = await session.execute(total_stmt)
|
||||||
|
total = total_result.scalar_one()
|
||||||
|
|
||||||
|
items_stmt = (
|
||||||
|
select(Notification)
|
||||||
|
.where(*where_clauses)
|
||||||
|
.order_by(Notification.created_at.desc())
|
||||||
|
.limit(limit)
|
||||||
|
.offset(offset)
|
||||||
|
)
|
||||||
|
items_result = await session.execute(items_stmt)
|
||||||
|
items = list(items_result.scalars().all())
|
||||||
|
|
||||||
|
return items, total
|
||||||
|
|
||||||
|
async def get_unread_count(
|
||||||
|
self,
|
||||||
|
session: AsyncSession,
|
||||||
|
user_id: uuid.UUID,
|
||||||
|
) -> int:
|
||||||
|
"""Count unread, non-dismissed notifications for a user.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
session: Database session.
|
||||||
|
user_id: Owner of the notifications.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Number of unread notifications.
|
||||||
|
"""
|
||||||
|
stmt = (
|
||||||
|
select(func.count())
|
||||||
|
.select_from(Notification)
|
||||||
|
.where(
|
||||||
|
Notification.user_id == user_id,
|
||||||
|
Notification.read_at.is_(None),
|
||||||
|
Notification.dismissed_at.is_(None),
|
||||||
|
)
|
||||||
|
)
|
||||||
|
result = await session.execute(stmt)
|
||||||
|
return result.scalar_one()
|
||||||
|
|
||||||
|
async def mark_read(
|
||||||
|
self,
|
||||||
|
session: AsyncSession,
|
||||||
|
notification_id: uuid.UUID,
|
||||||
|
user_id: uuid.UUID,
|
||||||
|
) -> Notification:
|
||||||
|
"""Mark a single notification as read.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
session: Database session.
|
||||||
|
notification_id: UUID of the notification to mark.
|
||||||
|
user_id: Owner of the notification.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
The updated Notification instance.
|
||||||
|
|
||||||
|
Raises:
|
||||||
|
ValueError: If the notification does not exist or is not owned by the user.
|
||||||
|
"""
|
||||||
|
notification = await self._get_owned_notification(
|
||||||
|
session, notification_id, user_id
|
||||||
|
)
|
||||||
|
notification.read_at = datetime.now(timezone.utc)
|
||||||
|
await session.commit()
|
||||||
|
await session.refresh(notification)
|
||||||
|
return notification
|
||||||
|
|
||||||
|
async def mark_all_read(
|
||||||
|
self,
|
||||||
|
session: AsyncSession,
|
||||||
|
user_id: uuid.UUID,
|
||||||
|
) -> int:
|
||||||
|
"""Mark all unread notifications as read for a user.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
session: Database session.
|
||||||
|
user_id: Owner of the notifications.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Number of rows updated.
|
||||||
|
"""
|
||||||
|
stmt = (
|
||||||
|
update(Notification)
|
||||||
|
.where(
|
||||||
|
Notification.user_id == user_id,
|
||||||
|
Notification.read_at.is_(None),
|
||||||
|
Notification.dismissed_at.is_(None),
|
||||||
|
)
|
||||||
|
.values(read_at=datetime.now(timezone.utc))
|
||||||
|
)
|
||||||
|
result: CursorResult[Any] = await session.execute(stmt) # type: ignore[assignment]
|
||||||
|
await session.commit()
|
||||||
|
return result.rowcount or 0
|
||||||
|
|
||||||
|
async def dismiss_all(
|
||||||
|
self,
|
||||||
|
session: AsyncSession,
|
||||||
|
user_id: uuid.UUID,
|
||||||
|
) -> int:
|
||||||
|
"""Soft-delete all non-dismissed notifications for a user.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
session: Database session.
|
||||||
|
user_id: Owner of the notifications.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Number of rows updated.
|
||||||
|
"""
|
||||||
|
stmt = (
|
||||||
|
update(Notification)
|
||||||
|
.where(
|
||||||
|
Notification.user_id == user_id,
|
||||||
|
Notification.dismissed_at.is_(None),
|
||||||
|
)
|
||||||
|
.values(dismissed_at=datetime.now(timezone.utc))
|
||||||
|
)
|
||||||
|
result: CursorResult[Any] = await session.execute(stmt) # type: ignore[assignment]
|
||||||
|
await session.commit()
|
||||||
|
return result.rowcount or 0
|
||||||
|
|
||||||
|
async def dismiss(
|
||||||
|
self,
|
||||||
|
session: AsyncSession,
|
||||||
|
notification_id: uuid.UUID,
|
||||||
|
user_id: uuid.UUID,
|
||||||
|
) -> None:
|
||||||
|
"""Soft-delete a notification by setting dismissed_at.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
session: Database session.
|
||||||
|
notification_id: UUID of the notification to dismiss.
|
||||||
|
user_id: Owner of the notification.
|
||||||
|
|
||||||
|
Raises:
|
||||||
|
ValueError: If the notification does not exist or is not owned by the user.
|
||||||
|
"""
|
||||||
|
notification = await self._get_owned_notification(
|
||||||
|
session, notification_id, user_id
|
||||||
|
)
|
||||||
|
notification.dismissed_at = datetime.now(timezone.utc)
|
||||||
|
await session.commit()
|
||||||
|
|
||||||
|
async def _get_owned_notification(
|
||||||
|
self,
|
||||||
|
session: AsyncSession,
|
||||||
|
notification_id: uuid.UUID,
|
||||||
|
user_id: uuid.UUID,
|
||||||
|
) -> Notification:
|
||||||
|
"""Fetch a notification and verify ownership.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
session: Database session.
|
||||||
|
notification_id: UUID of the notification.
|
||||||
|
user_id: Expected owner.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
The Notification instance.
|
||||||
|
|
||||||
|
Raises:
|
||||||
|
ValueError: If the notification does not exist or is not owned.
|
||||||
|
"""
|
||||||
|
notification = await session.get(Notification, notification_id)
|
||||||
|
if notification is None or notification.user_id != user_id:
|
||||||
|
raise ValueError("Notification not found")
|
||||||
|
return notification
|
||||||
|
|
||||||
|
|
||||||
|
# Module-level singleton instance
|
||||||
|
notification_service = NotificationService()
|
||||||
@@ -0,0 +1,310 @@
|
|||||||
|
"""Permission fixer: applies mount permission policies post-start."""
|
||||||
|
|
||||||
|
import logging
|
||||||
|
import subprocess
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
def apply_mount_permissions(
|
||||||
|
container_id: str,
|
||||||
|
mounts: list[dict],
|
||||||
|
timeout: int = 10,
|
||||||
|
) -> list[dict[str, Any]]:
|
||||||
|
"""Apply permission policies to mounted directories in a running container.
|
||||||
|
|
||||||
|
Runs `chown`, `chmod`, and file-mode fixes for each mount that declares
|
||||||
|
an owner, mode, or file_mode. Requires the container to have a root user.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
container_id: Docker container ID or name.
|
||||||
|
mounts: List of mount definitions from the manifest.
|
||||||
|
timeout: Max seconds per docker exec command.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
List of result dicts: [{mount_name, success, error}]
|
||||||
|
"""
|
||||||
|
results = []
|
||||||
|
|
||||||
|
for mount in mounts:
|
||||||
|
name = mount.get("name", "unknown")
|
||||||
|
target = mount["target"]
|
||||||
|
owner = mount.get("owner")
|
||||||
|
mode = mount.get("mode")
|
||||||
|
file_mode = mount.get("file_mode")
|
||||||
|
|
||||||
|
result: dict[str, Any] = {
|
||||||
|
"mount_name": name,
|
||||||
|
"success": True,
|
||||||
|
"error": None,
|
||||||
|
}
|
||||||
|
|
||||||
|
# Skip read-only mounts — their permissions cannot be changed
|
||||||
|
# post-start because the bind mount is locked.
|
||||||
|
if mount.get("readonly", False):
|
||||||
|
logger.debug(
|
||||||
|
"Skipping permission fix for read-only mount %s (target=%s)",
|
||||||
|
name,
|
||||||
|
target,
|
||||||
|
)
|
||||||
|
results.append(result)
|
||||||
|
continue
|
||||||
|
|
||||||
|
# Skip if no permission policy defined
|
||||||
|
if not owner and not mode and not file_mode:
|
||||||
|
results.append(result)
|
||||||
|
continue
|
||||||
|
|
||||||
|
try:
|
||||||
|
if owner:
|
||||||
|
_run_in_container(
|
||||||
|
container_id,
|
||||||
|
["chown", "-R", f"{owner}:{owner}", target],
|
||||||
|
timeout,
|
||||||
|
)
|
||||||
|
logger.debug(
|
||||||
|
"Applied owner %s to %s in container %s",
|
||||||
|
owner,
|
||||||
|
target,
|
||||||
|
container_id,
|
||||||
|
)
|
||||||
|
|
||||||
|
if mode and result["success"]:
|
||||||
|
_run_in_container(
|
||||||
|
container_id,
|
||||||
|
["chmod", mode, target],
|
||||||
|
timeout,
|
||||||
|
)
|
||||||
|
logger.debug(
|
||||||
|
"Applied mode %s to %s in container %s",
|
||||||
|
mode,
|
||||||
|
target,
|
||||||
|
container_id,
|
||||||
|
)
|
||||||
|
|
||||||
|
if file_mode and result["success"]:
|
||||||
|
_run_in_container(
|
||||||
|
container_id,
|
||||||
|
[
|
||||||
|
"sh",
|
||||||
|
"-c",
|
||||||
|
f"find {target} -type f -exec chmod {file_mode} {{}} +",
|
||||||
|
],
|
||||||
|
timeout,
|
||||||
|
)
|
||||||
|
logger.debug(
|
||||||
|
"Applied file_mode %s to files in %s in container %s",
|
||||||
|
file_mode,
|
||||||
|
target,
|
||||||
|
container_id,
|
||||||
|
)
|
||||||
|
|
||||||
|
except PermissionFixError as exc:
|
||||||
|
result["success"] = False
|
||||||
|
result["error"] = str(exc)
|
||||||
|
logger.warning(
|
||||||
|
"Permission fix failed for mount %s (target=%s): %s",
|
||||||
|
name,
|
||||||
|
target,
|
||||||
|
exc,
|
||||||
|
)
|
||||||
|
|
||||||
|
results.append(result)
|
||||||
|
|
||||||
|
return results
|
||||||
|
|
||||||
|
|
||||||
|
def _exec_and_log(
|
||||||
|
container_id: str,
|
||||||
|
command: list[str],
|
||||||
|
timeout: int,
|
||||||
|
description: str,
|
||||||
|
) -> str:
|
||||||
|
"""Run a docker exec command and log stdout/stderr for debugging."""
|
||||||
|
cmd = ["docker", "exec", "--user", "root", container_id] + command
|
||||||
|
logger.debug("[SSH-fix] %s: %s", description, " ".join(cmd))
|
||||||
|
|
||||||
|
try:
|
||||||
|
result = subprocess.run(
|
||||||
|
cmd,
|
||||||
|
capture_output=True,
|
||||||
|
text=True,
|
||||||
|
timeout=timeout,
|
||||||
|
)
|
||||||
|
except subprocess.TimeoutExpired:
|
||||||
|
raise PermissionFixError(
|
||||||
|
f"Command timed out after {timeout}s: {' '.join(command)}"
|
||||||
|
)
|
||||||
|
except FileNotFoundError:
|
||||||
|
raise PermissionFixError(f"Docker command not found: {' '.join(command)}")
|
||||||
|
|
||||||
|
stdout = result.stdout.strip()
|
||||||
|
stderr = result.stderr.strip()
|
||||||
|
if stdout:
|
||||||
|
logger.debug("[SSH-fix] %s stdout: %s", description, stdout)
|
||||||
|
if stderr:
|
||||||
|
logger.debug("[SSH-fix] %s stderr: %s", description, stderr)
|
||||||
|
|
||||||
|
if result.returncode != 0:
|
||||||
|
raise PermissionFixError(
|
||||||
|
f"Command failed (rc={result.returncode}): {stderr or '(no stderr)'}"
|
||||||
|
)
|
||||||
|
return stdout
|
||||||
|
|
||||||
|
|
||||||
|
def apply_ssh_permissions(
|
||||||
|
container_id: str,
|
||||||
|
ssh_target: str,
|
||||||
|
container_user: str,
|
||||||
|
timeout: int = 10,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
"""Fix SSH directory ownership and permissions in a running container.
|
||||||
|
|
||||||
|
Runs chown and chmod on the ~/.ssh directory so the container user
|
||||||
|
can use the keys (SSH requires the private key to be owned by the
|
||||||
|
user with mode 600).
|
||||||
|
|
||||||
|
Args:
|
||||||
|
container_id: Docker container ID or name.
|
||||||
|
ssh_target: Absolute path to the .ssh directory inside the container.
|
||||||
|
container_user: The container user that should own the keys.
|
||||||
|
timeout: Max seconds per docker exec command.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Result dict with keys: success, error.
|
||||||
|
"""
|
||||||
|
result: dict[str, Any] = {"success": True, "error": None}
|
||||||
|
try:
|
||||||
|
# 1. Ensure directory is owned by the container user
|
||||||
|
_exec_and_log(
|
||||||
|
container_id,
|
||||||
|
["chown", "-R", f"{container_user}:{container_user}", ssh_target],
|
||||||
|
timeout,
|
||||||
|
"chown",
|
||||||
|
)
|
||||||
|
|
||||||
|
# 2. Set directory permissions
|
||||||
|
_exec_and_log(
|
||||||
|
container_id,
|
||||||
|
["chmod", "700", ssh_target],
|
||||||
|
timeout,
|
||||||
|
"chmod-dir",
|
||||||
|
)
|
||||||
|
|
||||||
|
# 3. Set private key permissions (id_ed25519, id_rsa, etc.)
|
||||||
|
_exec_and_log(
|
||||||
|
container_id,
|
||||||
|
[
|
||||||
|
"sh",
|
||||||
|
"-c",
|
||||||
|
f"find {ssh_target} -name 'id_*' -type f -exec chmod 600 {{}} +",
|
||||||
|
],
|
||||||
|
timeout,
|
||||||
|
"chmod-keys",
|
||||||
|
)
|
||||||
|
|
||||||
|
# 4. Verify final state
|
||||||
|
ls_output = _exec_and_log(
|
||||||
|
container_id,
|
||||||
|
["ls", "-la", ssh_target],
|
||||||
|
timeout,
|
||||||
|
"verify-ls",
|
||||||
|
)
|
||||||
|
stat_output = _exec_and_log(
|
||||||
|
container_id,
|
||||||
|
["stat", "-c", "%U:%G %a %n", ssh_target],
|
||||||
|
timeout,
|
||||||
|
"verify-stat-dir",
|
||||||
|
)
|
||||||
|
key_stat = _exec_and_log(
|
||||||
|
container_id,
|
||||||
|
[
|
||||||
|
"sh",
|
||||||
|
"-c",
|
||||||
|
f"stat -c '%U:%G %a %n' {ssh_target}/id_* 2>/dev/null || echo 'no id_* files found'",
|
||||||
|
],
|
||||||
|
timeout,
|
||||||
|
"verify-stat-keys",
|
||||||
|
)
|
||||||
|
|
||||||
|
logger.info(
|
||||||
|
"SSH permissions fixed for container %s (user=%s, target=%s). "
|
||||||
|
"ls:\n%s\nstat-dir: %s\nstat-keys: %s",
|
||||||
|
container_id,
|
||||||
|
container_user,
|
||||||
|
ssh_target,
|
||||||
|
ls_output,
|
||||||
|
stat_output,
|
||||||
|
key_stat,
|
||||||
|
)
|
||||||
|
except PermissionFixError as exc:
|
||||||
|
result["success"] = False
|
||||||
|
result["error"] = str(exc)
|
||||||
|
logger.warning(
|
||||||
|
"SSH permission fix failed for container %s (target=%s): %s",
|
||||||
|
container_id,
|
||||||
|
ssh_target,
|
||||||
|
exc,
|
||||||
|
)
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
class PermissionFixError(Exception):
|
||||||
|
"""Raised when a permission fix command fails."""
|
||||||
|
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
|
def _run_in_container(
|
||||||
|
container_id: str,
|
||||||
|
command: list[str],
|
||||||
|
timeout: int,
|
||||||
|
) -> None:
|
||||||
|
"""Run a command inside a container as root.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
container_id: Docker container ID or name.
|
||||||
|
command: Command + args to execute.
|
||||||
|
timeout: Max seconds to wait.
|
||||||
|
|
||||||
|
Raises:
|
||||||
|
PermissionFixError: If the command fails or times out.
|
||||||
|
"""
|
||||||
|
cmd = ["docker", "exec", "--user", "root", container_id] + command
|
||||||
|
|
||||||
|
try:
|
||||||
|
result = subprocess.run(
|
||||||
|
cmd,
|
||||||
|
capture_output=True,
|
||||||
|
text=True,
|
||||||
|
timeout=timeout,
|
||||||
|
)
|
||||||
|
except subprocess.TimeoutExpired:
|
||||||
|
raise PermissionFixError(
|
||||||
|
f"Command timed out after {timeout}s: {' '.join(command)}"
|
||||||
|
)
|
||||||
|
except FileNotFoundError:
|
||||||
|
raise PermissionFixError(f"Docker command not found: {' '.join(command)}")
|
||||||
|
|
||||||
|
if result.returncode != 0:
|
||||||
|
raise PermissionFixError(
|
||||||
|
f"Command failed (rc={result.returncode}): {result.stderr.strip()}"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def check_root_user_available(container_id: str, timeout: int = 5) -> bool:
|
||||||
|
"""Check if the container has a root user we can exec as.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
container_id: Docker container ID or name.
|
||||||
|
timeout: Max seconds to wait.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
True if root user exists and is usable.
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
_run_in_container(container_id, ["id", "root"], timeout)
|
||||||
|
return True
|
||||||
|
except PermissionFixError:
|
||||||
|
return False
|
||||||
@@ -1,5 +1,6 @@
|
|||||||
"""SSH key service utilities for preparing keys for container use."""
|
"""SSH key service utilities for preparing keys for container use."""
|
||||||
|
|
||||||
|
import logging
|
||||||
import os
|
import os
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
@@ -7,6 +8,8 @@ from cryptography.fernet import Fernet
|
|||||||
|
|
||||||
from src.config import Settings
|
from src.config import Settings
|
||||||
|
|
||||||
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
def _get_fernet() -> Fernet:
|
def _get_fernet() -> Fernet:
|
||||||
"""Generate a valid Fernet key from the session secret."""
|
"""Generate a valid Fernet key from the session secret."""
|
||||||
@@ -19,17 +22,26 @@ def _get_fernet() -> Fernet:
|
|||||||
return Fernet(key)
|
return Fernet(key)
|
||||||
|
|
||||||
|
|
||||||
def prepare_ssh_key_files(instance_dir: str, ssh_key) -> str:
|
def prepare_ssh_key_files(
|
||||||
|
instance_dir: str,
|
||||||
|
ssh_key,
|
||||||
|
subdir: str = ".ssh",
|
||||||
|
uid: int | None = None,
|
||||||
|
gid: int | None = None,
|
||||||
|
) -> str:
|
||||||
"""Decrypt and write SSH key files to instance directory for container mounting.
|
"""Decrypt and write SSH key files to instance directory for container mounting.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
instance_dir: Path to instance directory
|
instance_dir: Path to instance directory
|
||||||
ssh_key: SSHKey model instance with encrypted private key
|
ssh_key: SSHKey model instance with encrypted private key
|
||||||
|
subdir: Subdirectory within instance_dir to write to (default: ".ssh")
|
||||||
|
uid: Optional UID to own the files (for bind-mount into non-root container)
|
||||||
|
gid: Optional GID to own the files
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
Path to the .ssh directory
|
Path to the .ssh directory
|
||||||
"""
|
"""
|
||||||
ssh_dir = Path(instance_dir) / ".ssh"
|
ssh_dir = Path(instance_dir) / subdir
|
||||||
ssh_dir.mkdir(parents=True, exist_ok=True)
|
ssh_dir.mkdir(parents=True, exist_ok=True)
|
||||||
|
|
||||||
# Decrypt private key
|
# Decrypt private key
|
||||||
@@ -57,6 +69,30 @@ def prepare_ssh_key_files(instance_dir: str, ssh_key) -> str:
|
|||||||
config_path.write_text(config_content)
|
config_path.write_text(config_content)
|
||||||
os.chmod(config_path, 0o644)
|
os.chmod(config_path, 0o644)
|
||||||
|
|
||||||
|
# Set ownership to target container user if requested
|
||||||
|
if uid is not None or gid is not None:
|
||||||
|
effective_uid = uid if uid is not None else -1
|
||||||
|
effective_gid = gid if gid is not None else -1
|
||||||
|
try:
|
||||||
|
os.chown(ssh_dir, effective_uid, effective_gid)
|
||||||
|
os.chown(private_key_path, effective_uid, effective_gid)
|
||||||
|
os.chown(public_key_path, effective_uid, effective_gid)
|
||||||
|
os.chown(config_path, effective_uid, effective_gid)
|
||||||
|
logger.debug(
|
||||||
|
"Set SSH key ownership to uid=%s gid=%s for %s",
|
||||||
|
effective_uid,
|
||||||
|
effective_gid,
|
||||||
|
ssh_dir,
|
||||||
|
)
|
||||||
|
except PermissionError as exc:
|
||||||
|
logger.warning(
|
||||||
|
"Cannot chown SSH keys to uid=%s gid=%s (running as uid=%s): %s",
|
||||||
|
effective_uid,
|
||||||
|
effective_gid,
|
||||||
|
os.getuid(),
|
||||||
|
exc,
|
||||||
|
)
|
||||||
|
|
||||||
return str(ssh_dir)
|
return str(ssh_dir)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -3,21 +3,37 @@
|
|||||||
import asyncio
|
import asyncio
|
||||||
import logging
|
import logging
|
||||||
import uuid
|
import uuid
|
||||||
from typing import Any
|
from datetime import datetime, timezone
|
||||||
|
|
||||||
from fastapi import WebSocket
|
from fastapi import WebSocket
|
||||||
|
|
||||||
|
from src.database import SessionLocal
|
||||||
|
from src.models.terminal_session import TerminalSessionModel
|
||||||
from src.services.terminal_session import TerminalSession
|
from src.services.terminal_session import TerminalSession
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
||||||
|
class MaxSessionsExceededError(Exception):
|
||||||
|
"""Raised when the maximum number of terminal sessions per instance is reached."""
|
||||||
|
|
||||||
|
def __init__(self, instance_id: str, max_sessions: int = 5) -> None:
|
||||||
|
self.instance_id = instance_id
|
||||||
|
self.max_sessions = max_sessions
|
||||||
|
super().__init__(
|
||||||
|
f"Maximum of {max_sessions} terminal sessions reached for instance {instance_id}"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
class TerminalManager:
|
class TerminalManager:
|
||||||
"""Manages active terminal sessions with persistence support."""
|
"""Manages active terminal sessions with persistence support."""
|
||||||
|
|
||||||
|
# Maximum sessions per tool instance
|
||||||
|
MAX_SESSIONS_PER_INSTANCE = 5
|
||||||
|
|
||||||
def __init__(self) -> None:
|
def __init__(self) -> None:
|
||||||
# Track sessions by instance_id for persistence
|
# Track sessions by (instance_id, session_id) for multi-session support
|
||||||
self._sessions: dict[str, TerminalSession] = {}
|
self._sessions: dict[tuple[str, str], TerminalSession] = {}
|
||||||
self._idle_check_task: asyncio.Task | None = None
|
self._idle_check_task: asyncio.Task | None = None
|
||||||
self._start_idle_check()
|
self._start_idle_check()
|
||||||
|
|
||||||
@@ -43,77 +59,280 @@ class TerminalManager:
|
|||||||
|
|
||||||
async def _cleanup_idle_sessions(self) -> None:
|
async def _cleanup_idle_sessions(self) -> None:
|
||||||
"""Clean up sessions that have been idle for too long."""
|
"""Clean up sessions that have been idle for too long."""
|
||||||
idle_sessions = []
|
idle_keys = []
|
||||||
for instance_id, session in list(self._sessions.items()):
|
for (instance_id, session_id), session in list(self._sessions.items()):
|
||||||
if session.is_idle():
|
if session.is_idle():
|
||||||
idle_sessions.append(instance_id)
|
idle_keys.append((instance_id, session_id))
|
||||||
|
|
||||||
for instance_id in idle_sessions:
|
for key in idle_keys:
|
||||||
logger.info("Cleaning up idle terminal session for instance %s", instance_id)
|
instance_id, session_id = key
|
||||||
session = self._sessions.pop(instance_id, None)
|
logger.info(
|
||||||
|
"Cleaning up idle terminal session %s for instance %s",
|
||||||
|
session_id,
|
||||||
|
instance_id,
|
||||||
|
)
|
||||||
|
session = self._sessions.pop(key, None)
|
||||||
if session:
|
if session:
|
||||||
await session.close()
|
await session.close()
|
||||||
|
# Update DB status fire-and-forget
|
||||||
|
asyncio.create_task(self._mark_closed_in_db(session_id))
|
||||||
|
|
||||||
|
async def _insert_db_session_row(
|
||||||
|
self,
|
||||||
|
session_id: str,
|
||||||
|
instance_id: uuid.UUID,
|
||||||
|
name: str,
|
||||||
|
) -> None:
|
||||||
|
"""Insert a TerminalSessionModel row into the database."""
|
||||||
|
try:
|
||||||
|
async with SessionLocal() as db_session:
|
||||||
|
db_row = TerminalSessionModel(
|
||||||
|
id=uuid.UUID(session_id),
|
||||||
|
instance_id=instance_id,
|
||||||
|
name=name,
|
||||||
|
status="active",
|
||||||
|
created_at=datetime.now(timezone.utc),
|
||||||
|
last_activity_at=datetime.now(timezone.utc),
|
||||||
|
)
|
||||||
|
db_session.add(db_row)
|
||||||
|
await db_session.commit()
|
||||||
|
logger.debug(
|
||||||
|
"Inserted terminal session row %s for instance %s",
|
||||||
|
session_id,
|
||||||
|
instance_id,
|
||||||
|
)
|
||||||
|
except Exception as exc:
|
||||||
|
logger.error("Failed to insert terminal session row: %s", exc)
|
||||||
|
|
||||||
|
async def _mark_closed_in_db(self, session_id: str) -> None:
|
||||||
|
"""Mark a terminal session as closed in the database."""
|
||||||
|
try:
|
||||||
|
async with SessionLocal() as db_session:
|
||||||
|
db_row = await db_session.get(
|
||||||
|
TerminalSessionModel, uuid.UUID(session_id)
|
||||||
|
)
|
||||||
|
if db_row:
|
||||||
|
db_row.status = "closed"
|
||||||
|
db_row.closed_at = datetime.now(timezone.utc)
|
||||||
|
await db_session.commit()
|
||||||
|
logger.debug(
|
||||||
|
"Marked terminal session %s as closed in DB", session_id
|
||||||
|
)
|
||||||
|
except Exception as exc:
|
||||||
|
logger.error("Failed to mark terminal session as closed in DB: %s", exc)
|
||||||
|
|
||||||
|
def _count_sessions_for_instance(self, instance_id_str: str) -> int:
|
||||||
|
"""Count active in-memory sessions for a given instance."""
|
||||||
|
return sum(1 for (iid, _sid) in self._sessions if iid == instance_id_str)
|
||||||
|
|
||||||
|
async def create_session(
|
||||||
|
self,
|
||||||
|
instance_id: uuid.UUID,
|
||||||
|
container_id: str,
|
||||||
|
startup_command: str | None = None,
|
||||||
|
name: str | None = None,
|
||||||
|
session_id: str | None = None,
|
||||||
|
) -> TerminalSession:
|
||||||
|
"""Create a new terminal session for an instance.
|
||||||
|
|
||||||
|
Enforces a maximum of MAX_SESSIONS_PER_INSTANCE sessions per instance.
|
||||||
|
Inserts a DB row fire-and-forget.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
instance_id: UUID of the tool instance.
|
||||||
|
container_id: Docker container ID.
|
||||||
|
startup_command: Optional startup command to run.
|
||||||
|
name: Optional session name (auto-generated if omitted).
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
The newly created TerminalSession.
|
||||||
|
|
||||||
|
Raises:
|
||||||
|
MaxSessionsExceededError: If the instance already has max sessions.
|
||||||
|
"""
|
||||||
|
instance_id_str = str(instance_id)
|
||||||
|
|
||||||
|
if (
|
||||||
|
self._count_sessions_for_instance(instance_id_str)
|
||||||
|
>= self.MAX_SESSIONS_PER_INSTANCE
|
||||||
|
):
|
||||||
|
raise MaxSessionsExceededError(
|
||||||
|
instance_id_str, self.MAX_SESSIONS_PER_INSTANCE
|
||||||
|
)
|
||||||
|
|
||||||
|
if session_id is None:
|
||||||
|
session_id = str(uuid.uuid4())
|
||||||
|
session = TerminalSession(
|
||||||
|
session_id=session_id,
|
||||||
|
instance_id=instance_id,
|
||||||
|
container_id=container_id,
|
||||||
|
startup_command=startup_command,
|
||||||
|
name=name,
|
||||||
|
)
|
||||||
|
await session.start(startup_command=startup_command)
|
||||||
|
|
||||||
|
key = (instance_id_str, session_id)
|
||||||
|
self._sessions[key] = session
|
||||||
|
|
||||||
|
# Fire-and-forget DB insert (skip if row already exists)
|
||||||
|
asyncio.create_task(
|
||||||
|
self._insert_db_session_row(session_id, instance_id, session.name)
|
||||||
|
)
|
||||||
|
|
||||||
|
logger.info(
|
||||||
|
"Created terminal session %s for instance %s (name=%s)",
|
||||||
|
session_id,
|
||||||
|
instance_id,
|
||||||
|
session.name,
|
||||||
|
)
|
||||||
|
return session
|
||||||
|
|
||||||
async def get_or_create_session(
|
async def get_or_create_session(
|
||||||
self,
|
self,
|
||||||
instance_id: uuid.UUID,
|
instance_id: uuid.UUID,
|
||||||
container_id: str,
|
container_id: str,
|
||||||
|
startup_command: str | None = None,
|
||||||
) -> TerminalSession:
|
) -> TerminalSession:
|
||||||
"""Get existing session or create a new one."""
|
"""Get existing session or create a new one.
|
||||||
|
|
||||||
|
Backward-compatible alias that uses 'default' as the session_id.
|
||||||
|
"""
|
||||||
# Ensure idle check is running (lazy start)
|
# Ensure idle check is running (lazy start)
|
||||||
self._start_idle_check()
|
self._start_idle_check()
|
||||||
|
|
||||||
instance_id_str = str(instance_id)
|
instance_id_str = str(instance_id)
|
||||||
|
key = (instance_id_str, "default")
|
||||||
# Check for existing session
|
|
||||||
if instance_id_str in self._sessions:
|
# Check for existing default session
|
||||||
session = self._sessions[instance_id_str]
|
if key in self._sessions:
|
||||||
|
session = self._sessions[key]
|
||||||
|
|
||||||
# Check if session is still alive
|
# Check if session is still alive
|
||||||
if session.is_alive():
|
if session.is_alive():
|
||||||
logger.info("Reattaching to existing terminal session for instance %s", instance_id)
|
logger.debug(
|
||||||
|
"Reattaching to existing terminal session for instance %s",
|
||||||
|
instance_id,
|
||||||
|
)
|
||||||
return session
|
return session
|
||||||
else:
|
else:
|
||||||
# Session died, clean it up
|
# Session died, clean it up
|
||||||
logger.info("Existing session for instance %s is dead, cleaning up", instance_id)
|
logger.debug(
|
||||||
|
"Existing session for instance %s is dead, cleaning up",
|
||||||
|
instance_id,
|
||||||
|
)
|
||||||
await session.close()
|
await session.close()
|
||||||
del self._sessions[instance_id_str]
|
del self._sessions[key]
|
||||||
|
|
||||||
# Create new session
|
# Create new default session
|
||||||
logger.info("Creating new terminal session for instance %s", instance_id)
|
logger.info(
|
||||||
|
"Creating new default terminal session for instance %s", instance_id
|
||||||
|
)
|
||||||
session_id = str(uuid.uuid4())
|
session_id = str(uuid.uuid4())
|
||||||
session = TerminalSession(session_id, instance_id, container_id)
|
session = TerminalSession(
|
||||||
await session.start()
|
session_id=session_id,
|
||||||
self._sessions[instance_id_str] = session
|
instance_id=instance_id,
|
||||||
|
container_id=container_id,
|
||||||
|
startup_command=startup_command,
|
||||||
|
name="Session 1",
|
||||||
|
)
|
||||||
|
await session.start(startup_command=startup_command)
|
||||||
|
self._sessions[key] = session
|
||||||
|
|
||||||
|
# Fire-and-forget DB insert
|
||||||
|
asyncio.create_task(
|
||||||
|
self._insert_db_session_row(session_id, instance_id, session.name)
|
||||||
|
)
|
||||||
|
|
||||||
return session
|
return session
|
||||||
|
|
||||||
|
def get_session(
|
||||||
|
self,
|
||||||
|
instance_id: str,
|
||||||
|
session_id: str,
|
||||||
|
) -> TerminalSession | None:
|
||||||
|
"""Lookup a session by composite key, or by internal session_id."""
|
||||||
|
session = self._sessions.get((instance_id, session_id))
|
||||||
|
if session is not None:
|
||||||
|
return session
|
||||||
|
# Fallback: search by internal TerminalSession.session_id
|
||||||
|
for (iid, _sid), sess in self._sessions.items():
|
||||||
|
if iid == instance_id and sess.session_id == session_id:
|
||||||
|
return sess
|
||||||
|
return None
|
||||||
|
|
||||||
|
def _find_key_by_internal_id(
|
||||||
|
self,
|
||||||
|
instance_id: str,
|
||||||
|
internal_session_id: str,
|
||||||
|
) -> tuple[str, str] | None:
|
||||||
|
"""Find the manager dict key for a session by its internal session_id."""
|
||||||
|
for (iid, sid), session in self._sessions.items():
|
||||||
|
if iid == instance_id and session.session_id == internal_session_id:
|
||||||
|
return (iid, sid)
|
||||||
|
return None
|
||||||
|
|
||||||
|
def get_sessions_for_instance(
|
||||||
|
self,
|
||||||
|
instance_id: str,
|
||||||
|
) -> list[TerminalSession]:
|
||||||
|
"""Return all in-memory sessions for a given instance."""
|
||||||
|
return [
|
||||||
|
session
|
||||||
|
for (iid, _sid), session in self._sessions.items()
|
||||||
|
if iid == instance_id
|
||||||
|
]
|
||||||
|
|
||||||
|
async def close_session(
|
||||||
|
self,
|
||||||
|
instance_id: str,
|
||||||
|
session_id: str,
|
||||||
|
) -> None:
|
||||||
|
"""Close a specific session and update its DB status."""
|
||||||
|
key = (instance_id, session_id)
|
||||||
|
session = self._sessions.pop(key, None)
|
||||||
|
if session:
|
||||||
|
await session.close()
|
||||||
|
# Fire-and-forget DB update
|
||||||
|
asyncio.create_task(self._mark_closed_in_db(session_id))
|
||||||
|
logger.info(
|
||||||
|
"Closed terminal session %s for instance %s",
|
||||||
|
session_id,
|
||||||
|
instance_id,
|
||||||
|
)
|
||||||
|
|
||||||
async def attach_websocket(
|
async def attach_websocket(
|
||||||
self,
|
self,
|
||||||
session: TerminalSession,
|
session: TerminalSession,
|
||||||
websocket: WebSocket,
|
websocket: WebSocket,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Attach a WebSocket to an existing session."""
|
"""Attach a WebSocket to an existing session.
|
||||||
# Handle concurrent connections - close existing ones
|
|
||||||
|
Closes existing WebSocket connections only for this specific session.
|
||||||
|
"""
|
||||||
|
# Handle concurrent connections - close existing ones within the same session
|
||||||
if session.has_websockets():
|
if session.has_websockets():
|
||||||
logger.info("Closing existing WebSocket connections for instance %s", session.instance_id)
|
logger.debug(
|
||||||
|
"Closing existing WebSocket connections for session %s (instance %s)",
|
||||||
|
session.session_id,
|
||||||
|
session.instance_id,
|
||||||
|
)
|
||||||
for ws in list(session._websockets):
|
for ws in list(session._websockets):
|
||||||
try:
|
try:
|
||||||
await ws.close(code=4000, reason="New connection established")
|
await ws.close(code=4000, reason="New connection established")
|
||||||
except Exception:
|
except Exception:
|
||||||
pass
|
pass # noqa: S110
|
||||||
session._websockets.clear()
|
session._websockets.clear()
|
||||||
|
|
||||||
# Attach new WebSocket
|
# Attach new WebSocket
|
||||||
session.attach_websocket(websocket)
|
session.attach_websocket(websocket)
|
||||||
|
|
||||||
# Replay buffer
|
# Replay buffer
|
||||||
buffer = session.get_buffer()
|
buffer = session.get_buffer()
|
||||||
if buffer:
|
if buffer:
|
||||||
try:
|
try:
|
||||||
await websocket.send_bytes(buffer)
|
await websocket.send_bytes(buffer)
|
||||||
except Exception:
|
except Exception:
|
||||||
pass
|
pass # noqa: S110
|
||||||
|
|
||||||
async def detach_websocket(
|
async def detach_websocket(
|
||||||
self,
|
self,
|
||||||
@@ -127,23 +346,61 @@ class TerminalManager:
|
|||||||
self,
|
self,
|
||||||
instance_id: uuid.UUID,
|
instance_id: uuid.UUID,
|
||||||
container_id: str,
|
container_id: str,
|
||||||
|
startup_command: str | None = None,
|
||||||
|
session_id: str | None = None,
|
||||||
|
name: str | None = None,
|
||||||
) -> TerminalSession:
|
) -> TerminalSession:
|
||||||
"""Reset a session by killing it and creating a new one."""
|
"""Reset a session by killing it and creating a new one.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
instance_id: UUID of the tool instance.
|
||||||
|
container_id: Docker container ID.
|
||||||
|
startup_command: Optional startup command.
|
||||||
|
session_id: Specific session to reset. If None, resets the default session.
|
||||||
|
name: Optional name to preserve for the new session.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
The newly created TerminalSession.
|
||||||
|
"""
|
||||||
instance_id_str = str(instance_id)
|
instance_id_str = str(instance_id)
|
||||||
|
target_session_id = session_id or "default"
|
||||||
|
key = (instance_id_str, target_session_id)
|
||||||
|
|
||||||
|
# Preserve old name if not provided
|
||||||
|
old_name = name
|
||||||
|
if old_name is None and key in self._sessions:
|
||||||
|
old_name = self._sessions[key].name
|
||||||
|
|
||||||
# Close existing session if any
|
# Close existing session if any
|
||||||
if instance_id_str in self._sessions:
|
if key in self._sessions:
|
||||||
logger.info("Resetting terminal session for instance %s", instance_id)
|
logger.debug(
|
||||||
old_session = self._sessions.pop(instance_id_str)
|
"Resetting terminal session %s for instance %s",
|
||||||
|
target_session_id,
|
||||||
|
instance_id,
|
||||||
|
)
|
||||||
|
old_session = self._sessions.pop(key)
|
||||||
await old_session.close()
|
await old_session.close()
|
||||||
|
# Fire-and-forget DB update for old session
|
||||||
# Create new session
|
asyncio.create_task(self._mark_closed_in_db(old_session.session_id))
|
||||||
session_id = str(uuid.uuid4())
|
|
||||||
session = TerminalSession(session_id, instance_id, container_id)
|
# Create new session preserving the same session_id slot
|
||||||
await session.start()
|
new_session_id = str(uuid.uuid4())
|
||||||
self._sessions[instance_id_str] = session
|
new_session = TerminalSession(
|
||||||
|
session_id=new_session_id,
|
||||||
return session
|
instance_id=instance_id,
|
||||||
|
container_id=container_id,
|
||||||
|
startup_command=startup_command,
|
||||||
|
name=old_name or ("Session 1" if target_session_id == "default" else None),
|
||||||
|
)
|
||||||
|
await new_session.start(startup_command=startup_command)
|
||||||
|
self._sessions[key] = new_session
|
||||||
|
|
||||||
|
# Fire-and-forget DB insert
|
||||||
|
asyncio.create_task(
|
||||||
|
self._insert_db_session_row(new_session_id, instance_id, new_session.name)
|
||||||
|
)
|
||||||
|
|
||||||
|
return new_session
|
||||||
|
|
||||||
async def close_all(self) -> None:
|
async def close_all(self) -> None:
|
||||||
"""Close all active sessions."""
|
"""Close all active sessions."""
|
||||||
@@ -151,7 +408,7 @@ class TerminalManager:
|
|||||||
self._sessions.clear()
|
self._sessions.clear()
|
||||||
for session in sessions:
|
for session in sessions:
|
||||||
await session.close()
|
await session.close()
|
||||||
|
|
||||||
if self._idle_check_task and not self._idle_check_task.done():
|
if self._idle_check_task and not self._idle_check_task.done():
|
||||||
self._idle_check_task.cancel()
|
self._idle_check_task.cancel()
|
||||||
|
|
||||||
|
|||||||
@@ -18,49 +18,82 @@ logger = logging.getLogger(__name__)
|
|||||||
|
|
||||||
class TerminalSession:
|
class TerminalSession:
|
||||||
"""Manages a single terminal session connected to a docker container.
|
"""Manages a single terminal session connected to a docker container.
|
||||||
|
|
||||||
Supports persistent sessions that survive WebSocket disconnections.
|
Supports persistent sessions that survive WebSocket disconnections.
|
||||||
Multiple WebSocket connections can attach/detach from the same session.
|
Multiple WebSocket connections can attach/detach from the same session.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
# Circular buffer size (10KB)
|
# Circular buffer size (10KB)
|
||||||
BUFFER_SIZE = 10 * 1024
|
BUFFER_SIZE = 10 * 1024
|
||||||
|
|
||||||
# Idle timeout in seconds (30 minutes)
|
# Idle timeout in seconds (30 minutes)
|
||||||
IDLE_TIMEOUT = 30 * 60
|
IDLE_TIMEOUT = 30 * 60
|
||||||
|
|
||||||
def __init__(self, session_id: str, instance_id: uuid.UUID, container_id: str) -> None:
|
# Session number counter per instance_id for auto-naming
|
||||||
|
_instance_counters: dict[str, int] = {}
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
session_id: str,
|
||||||
|
instance_id: uuid.UUID,
|
||||||
|
container_id: str,
|
||||||
|
startup_command: str | None = None,
|
||||||
|
name: str | None = None,
|
||||||
|
) -> None:
|
||||||
self.session_id = session_id
|
self.session_id = session_id
|
||||||
self.instance_id = instance_id
|
self.instance_id = instance_id
|
||||||
self.container_id = container_id
|
self.container_id = container_id
|
||||||
|
self.startup_command = startup_command
|
||||||
self.process: asyncio.subprocess.Process | None = None
|
self.process: asyncio.subprocess.Process | None = None
|
||||||
self._closed = False
|
self._closed = False
|
||||||
self._master_fd: int | None = None
|
self._master_fd: int | None = None
|
||||||
self._slave_fd: int | None = None
|
self._slave_fd: int | None = None
|
||||||
|
|
||||||
# Circular buffer for output replay
|
# Circular buffer for output replay
|
||||||
self._output_buffer: deque[bytes] = deque(maxlen=self.BUFFER_SIZE)
|
self._output_buffer: deque[bytes] = deque(maxlen=self.BUFFER_SIZE)
|
||||||
self._buffer_size = 0
|
self._buffer_size = 0
|
||||||
|
|
||||||
# WebSocket connections
|
# WebSocket connections
|
||||||
self._websockets: set[Any] = set()
|
self._websockets: set[Any] = set()
|
||||||
|
|
||||||
# Activity tracking
|
# Activity tracking
|
||||||
self.last_activity = time.time()
|
self.last_activity = time.time()
|
||||||
|
|
||||||
# Terminal size
|
# Terminal size
|
||||||
self._cols = 80
|
self._cols = 80
|
||||||
self._rows = 24
|
self._rows = 24
|
||||||
|
|
||||||
async def start(self) -> None:
|
# Session metadata
|
||||||
|
self.name = name or self._generate_name(str(instance_id))
|
||||||
|
self.status: str = "active"
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _generate_name(cls, instance_id: str) -> str:
|
||||||
|
"""Generate an auto-incremented session name for the instance."""
|
||||||
|
count = cls._instance_counters.get(instance_id, 0) + 1
|
||||||
|
cls._instance_counters[instance_id] = count
|
||||||
|
return f"Session {count}"
|
||||||
|
|
||||||
|
async def start(self, startup_command: str | None = None) -> None:
|
||||||
"""Start the docker exec process with a shell using a PTY."""
|
"""Start the docker exec process with a shell using a PTY."""
|
||||||
# Create a pseudo-terminal on the host
|
# Create a pseudo-terminal on the host
|
||||||
self._master_fd, self._slave_fd = pty.openpty()
|
self._master_fd, self._slave_fd = pty.openpty()
|
||||||
|
|
||||||
# Set the terminal size initially
|
# Set the terminal size initially
|
||||||
self._set_terminal_size(self._cols, self._rows)
|
self._set_terminal_size(self._cols, self._rows)
|
||||||
logger.info(f"Starting terminal session {self.session_id} for container {self.container_id} with initial size {self._cols}x{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
|
# Start docker exec with the slave fd as stdin/stdout/stderr
|
||||||
# Using -it because the slave fd IS a TTY
|
# Using -it because the slave fd IS a TTY
|
||||||
self.process = await asyncio.create_subprocess_exec(
|
self.process = await asyncio.create_subprocess_exec(
|
||||||
@@ -71,16 +104,17 @@ class TerminalSession:
|
|||||||
"TERM=xterm",
|
"TERM=xterm",
|
||||||
self.container_id,
|
self.container_id,
|
||||||
"bash",
|
"bash",
|
||||||
"-il",
|
"-c",
|
||||||
|
shell_cmd,
|
||||||
stdin=self._slave_fd,
|
stdin=self._slave_fd,
|
||||||
stdout=self._slave_fd,
|
stdout=self._slave_fd,
|
||||||
stderr=self._slave_fd,
|
stderr=self._slave_fd,
|
||||||
)
|
)
|
||||||
|
|
||||||
# Close slave fd in parent process
|
# Close slave fd in parent process
|
||||||
os.close(self._slave_fd)
|
os.close(self._slave_fd)
|
||||||
self._slave_fd = None
|
self._slave_fd = None
|
||||||
|
|
||||||
self.last_activity = time.time()
|
self.last_activity = time.time()
|
||||||
|
|
||||||
def _set_terminal_size(self, cols: int, rows: int) -> None:
|
def _set_terminal_size(self, cols: int, rows: int) -> None:
|
||||||
@@ -90,10 +124,10 @@ class TerminalSession:
|
|||||||
return
|
return
|
||||||
# TIOCSWINSZ = 0x5414 on Linux
|
# TIOCSWINSZ = 0x5414 on Linux
|
||||||
TIOCSWINSZ = 0x5414
|
TIOCSWINSZ = 0x5414
|
||||||
size = struct.pack('HHHH', rows, cols, 0, 0)
|
size = struct.pack("HHHH", rows, cols, 0, 0)
|
||||||
try:
|
try:
|
||||||
fcntl.ioctl(self._master_fd, TIOCSWINSZ, size)
|
fcntl.ioctl(self._master_fd, TIOCSWINSZ, size)
|
||||||
logger.info(f"Resized PTY to {cols}x{rows} (fd={self._master_fd})")
|
logger.debug(f"Resized PTY to {cols}x{rows} (fd={self._master_fd})")
|
||||||
except (OSError, IOError) as e:
|
except (OSError, IOError) as e:
|
||||||
logger.error(f"Failed to resize PTY: {e}")
|
logger.error(f"Failed to resize PTY: {e}")
|
||||||
|
|
||||||
@@ -118,7 +152,7 @@ class TerminalSession:
|
|||||||
"""Add data to circular buffer, maintaining size limit."""
|
"""Add data to circular buffer, maintaining size limit."""
|
||||||
self._output_buffer.append(data)
|
self._output_buffer.append(data)
|
||||||
self._buffer_size += len(data)
|
self._buffer_size += len(data)
|
||||||
|
|
||||||
# Trim if exceeds max size
|
# Trim if exceeds max size
|
||||||
while self._buffer_size > self.BUFFER_SIZE and self._output_buffer:
|
while self._buffer_size > self.BUFFER_SIZE and self._output_buffer:
|
||||||
removed = self._output_buffer.popleft()
|
removed = self._output_buffer.popleft()
|
||||||
@@ -143,16 +177,16 @@ class TerminalSession:
|
|||||||
if self._closed:
|
if self._closed:
|
||||||
logger.warning("Cannot resize: session is closed")
|
logger.warning("Cannot resize: session is closed")
|
||||||
return
|
return
|
||||||
|
|
||||||
# Only resize if dimensions actually changed
|
# Only resize if dimensions actually changed
|
||||||
if cols == self._cols and rows == self._rows:
|
if cols == self._cols and rows == self._rows:
|
||||||
return
|
return
|
||||||
|
|
||||||
self._cols = cols
|
self._cols = cols
|
||||||
self._rows = rows
|
self._rows = rows
|
||||||
logger.info(f"resize() called for session {self.session_id}: {cols}x{rows}")
|
logger.debug(f"resize() called for session {self.session_id}: {cols}x{rows}")
|
||||||
self._set_terminal_size(cols, rows)
|
self._set_terminal_size(cols, rows)
|
||||||
|
|
||||||
# Docker exec -it creates its own PTY inside the container,
|
# Docker exec -it creates its own PTY inside the container,
|
||||||
# so host PTY resize doesn't propagate to the container shell.
|
# so host PTY resize doesn't propagate to the container shell.
|
||||||
# Send SIGWINCH to the docker exec process on the host.
|
# Send SIGWINCH to the docker exec process on the host.
|
||||||
@@ -161,14 +195,19 @@ class TerminalSession:
|
|||||||
if self.process and self.process.pid:
|
if self.process and self.process.pid:
|
||||||
try:
|
try:
|
||||||
os.kill(self.process.pid, signal.SIGWINCH)
|
os.kill(self.process.pid, signal.SIGWINCH)
|
||||||
logger.debug(f"Sent SIGWINCH to docker exec process {self.process.pid} for session {self.session_id}")
|
logger.debug(
|
||||||
|
f"Sent SIGWINCH to docker exec process {self.process.pid} for session {self.session_id}"
|
||||||
|
)
|
||||||
except ProcessLookupError:
|
except ProcessLookupError:
|
||||||
logger.warning(f"docker exec process {self.process.pid} not found for session {self.session_id}")
|
logger.warning(
|
||||||
|
f"docker exec process {self.process.pid} not found for session {self.session_id}"
|
||||||
|
)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.warning(f"Failed to send SIGWINCH: {e}")
|
logger.warning(f"Failed to send SIGWINCH: {e}")
|
||||||
|
|
||||||
async def reset(self) -> None:
|
async def reset(self) -> None:
|
||||||
"""Reset the session by killing the process and clearing state."""
|
"""Reset the session by killing the process and clearing state."""
|
||||||
|
self.status = "resetting"
|
||||||
await self.close()
|
await self.close()
|
||||||
self._closed = False
|
self._closed = False
|
||||||
self._output_buffer.clear()
|
self._output_buffer.clear()
|
||||||
@@ -177,18 +216,20 @@ class TerminalSession:
|
|||||||
self.process = None
|
self.process = None
|
||||||
self._master_fd = None
|
self._master_fd = None
|
||||||
self._slave_fd = None
|
self._slave_fd = None
|
||||||
|
self.status = "active"
|
||||||
|
|
||||||
async def close(self) -> None:
|
async def close(self) -> None:
|
||||||
"""Close the session and cleanup."""
|
"""Close the session and cleanup."""
|
||||||
if self._closed:
|
if self._closed:
|
||||||
return
|
return
|
||||||
self._closed = True
|
self._closed = True
|
||||||
|
self.status = "closed"
|
||||||
|
|
||||||
if self._master_fd is not None:
|
if self._master_fd is not None:
|
||||||
try:
|
try:
|
||||||
os.close(self._master_fd)
|
os.close(self._master_fd)
|
||||||
except OSError:
|
except OSError:
|
||||||
pass
|
pass # noqa: S110
|
||||||
self._master_fd = None
|
self._master_fd = None
|
||||||
|
|
||||||
if self.process is not None:
|
if self.process is not None:
|
||||||
@@ -231,7 +272,7 @@ class TerminalSession:
|
|||||||
await ws.send_bytes(data)
|
await ws.send_bytes(data)
|
||||||
except Exception:
|
except Exception:
|
||||||
dead_sockets.add(ws)
|
dead_sockets.add(ws)
|
||||||
|
|
||||||
# Clean up dead sockets
|
# Clean up dead sockets
|
||||||
for ws in dead_sockets:
|
for ws in dead_sockets:
|
||||||
self._websockets.discard(ws)
|
self._websockets.discard(ws)
|
||||||
|
|||||||
@@ -2,7 +2,6 @@
|
|||||||
|
|
||||||
import subprocess
|
import subprocess
|
||||||
from dataclasses import dataclass, field
|
from dataclasses import dataclass, field
|
||||||
from typing import Any
|
|
||||||
|
|
||||||
|
|
||||||
def _run_git_command(repo_path: str, *args: str) -> str:
|
def _run_git_command(repo_path: str, *args: str) -> str:
|
||||||
|
|||||||
@@ -0,0 +1,67 @@
|
|||||||
|
"""Integration tests for multi-session terminal WebSocket and REST API."""
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
from fastapi.testclient import TestClient
|
||||||
|
from src.main import app
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def client():
|
||||||
|
return TestClient(app)
|
||||||
|
|
||||||
|
|
||||||
|
class TestTerminalWebSocketMultiSession:
|
||||||
|
"""Tests for multi-session WebSocket routing."""
|
||||||
|
|
||||||
|
def test_specific_session_websocket_route_exists(self, client):
|
||||||
|
"""The specific session WebSocket route should be registered."""
|
||||||
|
# We can't easily test WebSocket without auth, but we can verify
|
||||||
|
# the route exists by checking for a 403 (no auth cookie)
|
||||||
|
response = client.get("/ws/tool-instances/test-instance/terminal/test-session")
|
||||||
|
# WebSocket endpoint returns 403 when accessed via HTTP GET
|
||||||
|
assert response.status_code in (403, 404)
|
||||||
|
|
||||||
|
def test_default_session_alias_route_exists(self, client):
|
||||||
|
"""The default session alias route should still exist."""
|
||||||
|
response = client.get("/ws/tool-instances/test-instance/terminal")
|
||||||
|
assert response.status_code in (403, 404)
|
||||||
|
|
||||||
|
|
||||||
|
class TestTerminalRestApi:
|
||||||
|
"""Tests for REST API endpoints."""
|
||||||
|
|
||||||
|
def test_list_sessions_requires_auth(self, client):
|
||||||
|
"""List sessions endpoint requires authentication."""
|
||||||
|
response = client.get("/instances/test/terminal/sessions")
|
||||||
|
assert response.status_code == 401
|
||||||
|
|
||||||
|
def test_create_session_requires_auth(self, client):
|
||||||
|
"""Create session endpoint requires authentication."""
|
||||||
|
response = client.post(
|
||||||
|
"/instances/test/terminal/sessions",
|
||||||
|
json={},
|
||||||
|
)
|
||||||
|
assert response.status_code == 401
|
||||||
|
|
||||||
|
def test_close_session_requires_auth(self, client):
|
||||||
|
"""Close session endpoint requires authentication."""
|
||||||
|
response = client.delete("/instances/test/terminal/sessions/test-session")
|
||||||
|
assert response.status_code == 401
|
||||||
|
|
||||||
|
def test_reset_session_requires_auth(self, client):
|
||||||
|
"""Reset session endpoint requires authentication."""
|
||||||
|
response = client.post("/instances/test/terminal/sessions/test-session/reset")
|
||||||
|
assert response.status_code == 401
|
||||||
|
|
||||||
|
def test_rename_session_requires_auth(self, client):
|
||||||
|
"""Rename session endpoint requires authentication."""
|
||||||
|
response = client.post(
|
||||||
|
"/instances/test/terminal/sessions/test-session/rename",
|
||||||
|
json={"name": "New Name"},
|
||||||
|
)
|
||||||
|
assert response.status_code == 401
|
||||||
|
|
||||||
|
def test_legacy_reset_alias_requires_auth(self, client):
|
||||||
|
"""Legacy reset endpoint still requires auth."""
|
||||||
|
response = client.post("/instances/test/terminal/reset")
|
||||||
|
assert response.status_code == 401
|
||||||
@@ -8,16 +8,14 @@ from unittest.mock import patch
|
|||||||
import pytest
|
import pytest
|
||||||
import pytest_asyncio
|
import pytest_asyncio
|
||||||
from fastapi.testclient import TestClient
|
from fastapi.testclient import TestClient
|
||||||
from sqlalchemy import create_engine, text
|
|
||||||
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine
|
from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine
|
||||||
from sqlalchemy.orm import sessionmaker
|
|
||||||
|
|
||||||
# Set test environment BEFORE importing app modules
|
# Set test environment BEFORE importing app modules
|
||||||
os.environ["APP_ENV"] = "testing"
|
os.environ["APP_ENV"] = "testing"
|
||||||
os.environ["SECRET_KEY"] = "test-secret-key-for-testing-only-do-not-use-in-production"
|
os.environ["SECRET_KEY"] = "test-secret-key-for-testing-only-do-not-use-in-production"
|
||||||
os.environ["DATABASE_URL"] = "sqlite+aiosqlite:///:memory:"
|
os.environ["DATABASE_URL"] = "sqlite+aiosqlite:///:memory:"
|
||||||
|
|
||||||
from src.config import Settings, build_database_url
|
from src.config import Settings
|
||||||
from src.models.base import Base
|
from src.models.base import Base
|
||||||
from src.main import app
|
from src.main import app
|
||||||
from src.auth.dependencies import get_db_session
|
from src.auth.dependencies import get_db_session
|
||||||
@@ -133,6 +131,65 @@ def authenticated_client(test_client) -> Generator[TestClient, None, None]:
|
|||||||
yield test_client
|
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
|
@pytest.fixture
|
||||||
def admin_client(test_client) -> Generator[TestClient, None, None]:
|
def admin_client(test_client) -> Generator[TestClient, None, None]:
|
||||||
"""Provide an authenticated test client with an admin user."""
|
"""Provide an authenticated test client with an admin user."""
|
||||||
|
|||||||
@@ -1,255 +0,0 @@
|
|||||||
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"]
|
|
||||||
@@ -320,3 +320,134 @@ class TestConfigProfilesAPI:
|
|||||||
assert response.status_code == 200
|
assert response.status_code == 200
|
||||||
data = response.json()
|
data = response.json()
|
||||||
assert data["profile_id"] is None
|
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,268 @@
|
|||||||
|
"""Integration tests for SSE endpoint and lifecycle event flow."""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import uuid
|
||||||
|
from collections.abc import Generator
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
from fastapi.testclient import TestClient
|
||||||
|
from sqlalchemy import select
|
||||||
|
from sqlalchemy.ext.asyncio import AsyncSession
|
||||||
|
|
||||||
|
from src.api import events as events_module
|
||||||
|
from src.auth.session import decode_session_cookie
|
||||||
|
from src.config import Settings
|
||||||
|
from src.models.git_repository import GitRepository
|
||||||
|
from src.models.instance_event import InstanceEvent
|
||||||
|
from src.models.project import Project
|
||||||
|
from src.models.tool_instance import ToolInstance
|
||||||
|
from src.models.tool_type import ToolType
|
||||||
|
from src.services.event_bus import InstanceEventBus, InstanceEventPayload
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def event_bus() -> Generator[InstanceEventBus, None, None]:
|
||||||
|
bus = InstanceEventBus()
|
||||||
|
bus._reset_for_testing()
|
||||||
|
yield bus
|
||||||
|
bus._reset_for_testing()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def sample_payload() -> InstanceEventPayload:
|
||||||
|
return {
|
||||||
|
"event": "instance.started",
|
||||||
|
"instance_id": str(uuid.uuid4()),
|
||||||
|
"status": "starting",
|
||||||
|
"message": "Container starting...",
|
||||||
|
"metadata": {},
|
||||||
|
"timestamp": "2026-05-28T12:00:00Z",
|
||||||
|
"correlation_id": str(uuid.uuid4()),
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def _get_user_id_from_client(client: TestClient) -> uuid.UUID | None:
|
||||||
|
settings = Settings()
|
||||||
|
cookie = client.cookies.get("session")
|
||||||
|
if not cookie:
|
||||||
|
return None
|
||||||
|
session = decode_session_cookie(settings=settings, cookie_value=cookie)
|
||||||
|
if session and "user_id" in session:
|
||||||
|
return uuid.UUID(session["user_id"])
|
||||||
|
return None
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
def test_sse_requires_auth(test_client: TestClient) -> None:
|
||||||
|
response = test_client.get("/events/stream")
|
||||||
|
assert response.status_code == 401
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
def test_sse_enforces_connection_limit(authenticated_client: TestClient) -> None:
|
||||||
|
user_id = _get_user_id_from_client(authenticated_client)
|
||||||
|
assert user_id is not None
|
||||||
|
|
||||||
|
events_module._connection_counts[user_id] = events_module.MAX_CONNECTIONS_PER_USER
|
||||||
|
try:
|
||||||
|
response = authenticated_client.get("/events/stream")
|
||||||
|
assert response.status_code == 429
|
||||||
|
finally:
|
||||||
|
events_module._connection_counts.pop(user_id, None)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
def test_sse_event_generator_format() -> None:
|
||||||
|
"""Test the SSE endpoint is registered."""
|
||||||
|
from src.api.events import router
|
||||||
|
|
||||||
|
route_paths = [getattr(r, "path", "") for r in router.routes]
|
||||||
|
assert any("/stream" in str(p) for p in route_paths)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.integration
|
||||||
|
async def test_lifecycle_hook_publishes_event_and_persists(
|
||||||
|
authenticated_client: TestClient,
|
||||||
|
db_session: AsyncSession,
|
||||||
|
event_bus: InstanceEventBus,
|
||||||
|
) -> None:
|
||||||
|
"""Test that the lifecycle hook publishes an event and persists an audit row."""
|
||||||
|
user_id = _get_user_id_from_client(authenticated_client)
|
||||||
|
assert user_id is not None
|
||||||
|
|
||||||
|
project = Project(
|
||||||
|
id=uuid.uuid4(),
|
||||||
|
name="test-project",
|
||||||
|
description="Test",
|
||||||
|
owner_id=user_id,
|
||||||
|
)
|
||||||
|
repo = GitRepository(
|
||||||
|
id=uuid.uuid4(),
|
||||||
|
name="test-repo",
|
||||||
|
path="/tmp/test-repo",
|
||||||
|
project_id=project.id,
|
||||||
|
owner_id=user_id,
|
||||||
|
remote_url="https://github.com/test/repo.git",
|
||||||
|
)
|
||||||
|
tool_type = ToolType(
|
||||||
|
id=uuid.uuid4(),
|
||||||
|
name="test-tool",
|
||||||
|
display_name="Test Tool",
|
||||||
|
category="other",
|
||||||
|
interface_type="web",
|
||||||
|
requires_port=True,
|
||||||
|
default_port=8080,
|
||||||
|
definition_type="legacy",
|
||||||
|
compose_template="version: '3.8'\nservices:\n app:\n image: alpine\n command: sleep 3600\n",
|
||||||
|
)
|
||||||
|
db_session.add_all([project, repo, tool_type])
|
||||||
|
await db_session.commit()
|
||||||
|
|
||||||
|
instance = ToolInstance(
|
||||||
|
id=uuid.uuid4(),
|
||||||
|
name="test-instance",
|
||||||
|
display_name="Test Instance",
|
||||||
|
tool_type_id=tool_type.id,
|
||||||
|
repository_id=repo.id,
|
||||||
|
project_id=project.id,
|
||||||
|
owner_id=user_id,
|
||||||
|
status="pending",
|
||||||
|
compose_path="/tmp/test-compose.yml",
|
||||||
|
port=8080,
|
||||||
|
)
|
||||||
|
db_session.add(instance)
|
||||||
|
await db_session.commit()
|
||||||
|
|
||||||
|
received: list[Any] = []
|
||||||
|
|
||||||
|
def subscriber(payload: InstanceEventPayload) -> None:
|
||||||
|
received.append(payload)
|
||||||
|
|
||||||
|
event_bus.subscribe("instance.created", subscriber)
|
||||||
|
|
||||||
|
from src.services.lifecycle_hooks import publish_lifecycle_event
|
||||||
|
|
||||||
|
await publish_lifecycle_event(
|
||||||
|
event_bus=event_bus,
|
||||||
|
session=db_session,
|
||||||
|
instance=instance,
|
||||||
|
event_type="instance.created",
|
||||||
|
created_by=user_id,
|
||||||
|
status="pending",
|
||||||
|
message="Instance created",
|
||||||
|
)
|
||||||
|
|
||||||
|
assert len(received) == 1
|
||||||
|
assert received[0]["event"] == "instance.created"
|
||||||
|
|
||||||
|
result = await db_session.execute(
|
||||||
|
select(InstanceEvent).where(InstanceEvent.instance_id == instance.id)
|
||||||
|
)
|
||||||
|
rows = result.scalars().all()
|
||||||
|
assert len(rows) == 1
|
||||||
|
assert rows[0].event_type == "created"
|
||||||
|
assert rows[0].created_by == user_id
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.integration
|
||||||
|
async def test_lifecycle_event_persists_audit_row(
|
||||||
|
authenticated_client: TestClient,
|
||||||
|
db_session: AsyncSession,
|
||||||
|
event_bus: InstanceEventBus,
|
||||||
|
) -> None:
|
||||||
|
"""Test that publishing a lifecycle event persists an audit row."""
|
||||||
|
user_id = _get_user_id_from_client(authenticated_client)
|
||||||
|
assert user_id is not None
|
||||||
|
|
||||||
|
project = Project(
|
||||||
|
id=uuid.uuid4(),
|
||||||
|
name="test-project",
|
||||||
|
description="Test",
|
||||||
|
owner_id=user_id,
|
||||||
|
)
|
||||||
|
repo = GitRepository(
|
||||||
|
id=uuid.uuid4(),
|
||||||
|
name="test-repo",
|
||||||
|
path="/tmp/test-repo",
|
||||||
|
project_id=project.id,
|
||||||
|
owner_id=user_id,
|
||||||
|
remote_url="https://github.com/test/repo.git",
|
||||||
|
)
|
||||||
|
tool_type = ToolType(
|
||||||
|
id=uuid.uuid4(),
|
||||||
|
name="test-tool-2",
|
||||||
|
display_name="Test Tool 2",
|
||||||
|
category="other",
|
||||||
|
interface_type="web",
|
||||||
|
requires_port=True,
|
||||||
|
default_port=8080,
|
||||||
|
definition_type="legacy",
|
||||||
|
compose_template="version: '3.8'\nservices:\n app:\n image: alpine\n command: sleep 3600\n",
|
||||||
|
)
|
||||||
|
db_session.add_all([project, repo, tool_type])
|
||||||
|
await db_session.commit()
|
||||||
|
|
||||||
|
instance = ToolInstance(
|
||||||
|
id=uuid.uuid4(),
|
||||||
|
name="test-instance",
|
||||||
|
display_name="Test Instance",
|
||||||
|
tool_type_id=tool_type.id,
|
||||||
|
repository_id=repo.id,
|
||||||
|
project_id=project.id,
|
||||||
|
owner_id=user_id,
|
||||||
|
status="running",
|
||||||
|
compose_path="/tmp/test-compose.yml",
|
||||||
|
port=8080,
|
||||||
|
)
|
||||||
|
db_session.add(instance)
|
||||||
|
await db_session.commit()
|
||||||
|
|
||||||
|
from src.services.lifecycle_hooks import publish_lifecycle_event
|
||||||
|
|
||||||
|
await publish_lifecycle_event(
|
||||||
|
event_bus=event_bus,
|
||||||
|
session=db_session,
|
||||||
|
instance=instance,
|
||||||
|
event_type="instance.stopped",
|
||||||
|
created_by=user_id,
|
||||||
|
status="stopped",
|
||||||
|
message="Instance stopped",
|
||||||
|
)
|
||||||
|
|
||||||
|
result = await db_session.execute(
|
||||||
|
select(InstanceEvent).where(InstanceEvent.instance_id == instance.id)
|
||||||
|
)
|
||||||
|
rows = result.scalars().all()
|
||||||
|
assert len(rows) == 1
|
||||||
|
assert rows[0].event_type == "stopped"
|
||||||
|
assert rows[0].status == "stopped"
|
||||||
|
assert rows[0].created_by == user_id
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
def test_event_bus_pubsub(event_bus: InstanceEventBus) -> None:
|
||||||
|
"""Test that the event bus delivers events to subscribers."""
|
||||||
|
received: list[InstanceEventPayload] = []
|
||||||
|
|
||||||
|
def handler(payload: InstanceEventPayload) -> None:
|
||||||
|
received.append(payload)
|
||||||
|
|
||||||
|
event_bus.subscribe("test.event", handler)
|
||||||
|
|
||||||
|
payload: InstanceEventPayload = {
|
||||||
|
"event": "test.event",
|
||||||
|
"instance_id": str(uuid.uuid4()),
|
||||||
|
"status": "running",
|
||||||
|
"message": "Test",
|
||||||
|
"metadata": {},
|
||||||
|
"timestamp": "2026-05-28T12:00:00Z",
|
||||||
|
"correlation_id": str(uuid.uuid4()),
|
||||||
|
}
|
||||||
|
|
||||||
|
asyncio.run(event_bus.publish("test.event", payload))
|
||||||
|
|
||||||
|
assert len(received) == 1
|
||||||
|
assert received[0]["event"] == "test.event"
|
||||||
@@ -16,7 +16,6 @@ def test_base_metadata_collects_declared_tables() -> None:
|
|||||||
|
|
||||||
|
|
||||||
@pytest.mark.integration
|
@pytest.mark.integration
|
||||||
|
|
||||||
def test_shared_mixins_define_expected_columns() -> None:
|
def test_shared_mixins_define_expected_columns() -> None:
|
||||||
assert "id" in UUIDPrimaryKeyMixin.__dict__
|
assert "id" in UUIDPrimaryKeyMixin.__dict__
|
||||||
assert "created_at" in TimestampMixin.__dict__
|
assert "created_at" in TimestampMixin.__dict__
|
||||||
@@ -24,20 +23,26 @@ def test_shared_mixins_define_expected_columns() -> None:
|
|||||||
|
|
||||||
|
|
||||||
@pytest.mark.integration
|
@pytest.mark.integration
|
||||||
|
|
||||||
def test_expected_tables_are_registered() -> None:
|
def test_expected_tables_are_registered() -> None:
|
||||||
assert set(Base.metadata.tables) == {
|
assert set(Base.metadata.tables) == {
|
||||||
"refresh_tokens",
|
"config_profile_includes",
|
||||||
|
"config_profiles",
|
||||||
"git_repositories",
|
"git_repositories",
|
||||||
|
"health_checks",
|
||||||
|
"instance_events",
|
||||||
|
"notifications",
|
||||||
"projects",
|
"projects",
|
||||||
"ssh_keys",
|
"ssh_keys",
|
||||||
|
"terminal_sessions",
|
||||||
|
"tool_definition_manifests",
|
||||||
|
"tool_instances",
|
||||||
|
"tool_types",
|
||||||
"user_configs",
|
"user_configs",
|
||||||
"users",
|
"users",
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.integration
|
@pytest.mark.integration
|
||||||
|
|
||||||
def test_user_table_has_required_columns() -> None:
|
def test_user_table_has_required_columns() -> None:
|
||||||
columns = User.__table__.columns
|
columns = User.__table__.columns
|
||||||
|
|
||||||
@@ -56,7 +61,6 @@ def test_user_table_has_required_columns() -> None:
|
|||||||
|
|
||||||
|
|
||||||
@pytest.mark.integration
|
@pytest.mark.integration
|
||||||
|
|
||||||
def test_project_relationships_point_to_owner_and_default_ssh_key() -> None:
|
def test_project_relationships_point_to_owner_and_default_ssh_key() -> None:
|
||||||
owner_fk = next(iter(Project.__table__.c.owner_id.foreign_keys))
|
owner_fk = next(iter(Project.__table__.c.owner_id.foreign_keys))
|
||||||
ssh_fk = next(iter(Project.__table__.c.default_ssh_key_id.foreign_keys))
|
ssh_fk = next(iter(Project.__table__.c.default_ssh_key_id.foreign_keys))
|
||||||
@@ -68,7 +72,6 @@ def test_project_relationships_point_to_owner_and_default_ssh_key() -> None:
|
|||||||
|
|
||||||
|
|
||||||
@pytest.mark.integration
|
@pytest.mark.integration
|
||||||
|
|
||||||
def test_repository_and_user_config_relationships_are_registered() -> None:
|
def test_repository_and_user_config_relationships_are_registered() -> None:
|
||||||
project_fk = next(iter(GitRepository.__table__.c.project_id.foreign_keys))
|
project_fk = next(iter(GitRepository.__table__.c.project_id.foreign_keys))
|
||||||
owner_fk = next(iter(GitRepository.__table__.c.owner_id.foreign_keys))
|
owner_fk = next(iter(GitRepository.__table__.c.owner_id.foreign_keys))
|
||||||
@@ -82,33 +85,15 @@ def test_repository_and_user_config_relationships_are_registered() -> None:
|
|||||||
assert UserConfig.user.property.mapper.class_ is User
|
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.asyncio
|
||||||
@pytest.mark.integration
|
@pytest.mark.integration
|
||||||
|
|
||||||
async def test_async_session_can_insert_and_load_user(db_session: AsyncSession) -> None:
|
async def test_async_session_can_insert_and_load_user(db_session: AsyncSession) -> None:
|
||||||
user = User(email="dev@headquarter.local", name="Dev User", authentik_id="dev-user", avatar_url=None)
|
user = User(
|
||||||
|
email="dev@headquarter.local",
|
||||||
|
name="Dev User",
|
||||||
|
authentik_id="dev-user",
|
||||||
|
avatar_url=None,
|
||||||
|
)
|
||||||
|
|
||||||
db_session.add(user)
|
db_session.add(user)
|
||||||
await db_session.commit()
|
await db_session.commit()
|
||||||
|
|||||||
@@ -0,0 +1,326 @@
|
|||||||
|
"""Integration tests for notifications API."""
|
||||||
|
|
||||||
|
import uuid
|
||||||
|
from datetime import datetime, timedelta, timezone
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
from fastapi.testclient import TestClient
|
||||||
|
from sqlalchemy.ext.asyncio import AsyncSession
|
||||||
|
|
||||||
|
from src.models.user import User
|
||||||
|
from src.models.user_config import UserConfig
|
||||||
|
from src.services.notification_service import NotificationService
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def notification_service() -> NotificationService:
|
||||||
|
return NotificationService()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
async def user_a(db_session: AsyncSession) -> User:
|
||||||
|
user = User(
|
||||||
|
id=uuid.uuid4(),
|
||||||
|
email="user-a@headquarter.local",
|
||||||
|
name="User A",
|
||||||
|
authentik_id=f"authentik-{uuid.uuid4()}",
|
||||||
|
avatar_url=None,
|
||||||
|
)
|
||||||
|
db_session.add(user)
|
||||||
|
await db_session.commit()
|
||||||
|
return user
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
async def user_b(db_session: AsyncSession) -> User:
|
||||||
|
user = User(
|
||||||
|
id=uuid.uuid4(),
|
||||||
|
email="user-b@headquarter.local",
|
||||||
|
name="User B",
|
||||||
|
authentik_id=f"authentik-{uuid.uuid4()}",
|
||||||
|
avatar_url=None,
|
||||||
|
)
|
||||||
|
db_session.add(user)
|
||||||
|
await db_session.commit()
|
||||||
|
return user
|
||||||
|
|
||||||
|
|
||||||
|
def _mint_cookie_for_user(test_client: TestClient, user_id: uuid.UUID) -> None:
|
||||||
|
from src.auth.session import create_session_cookie
|
||||||
|
from src.config import Settings
|
||||||
|
|
||||||
|
settings = Settings()
|
||||||
|
cookie = create_session_cookie(
|
||||||
|
settings=settings,
|
||||||
|
user_id=str(user_id),
|
||||||
|
)
|
||||||
|
test_client.cookies.set("session", cookie)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
def test_list_requires_auth(test_client: TestClient) -> None:
|
||||||
|
response = test_client.get("/notifications")
|
||||||
|
assert response.status_code == 401
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
def test_list_returns_only_own_notifications(
|
||||||
|
authenticated_client: TestClient,
|
||||||
|
db_session: AsyncSession,
|
||||||
|
notification_service: NotificationService,
|
||||||
|
user_a: User,
|
||||||
|
user_b: User,
|
||||||
|
) -> None:
|
||||||
|
async def create_notifications() -> None:
|
||||||
|
await notification_service.create_notification(
|
||||||
|
db_session, user_a.id, category="instance", severity="info", title="A"
|
||||||
|
)
|
||||||
|
await notification_service.create_notification(
|
||||||
|
db_session, user_b.id, category="instance", severity="info", title="B"
|
||||||
|
)
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
|
||||||
|
asyncio.run(create_notifications())
|
||||||
|
|
||||||
|
_mint_cookie_for_user(authenticated_client, user_a.id)
|
||||||
|
response = authenticated_client.get("/notifications")
|
||||||
|
assert response.status_code == 200
|
||||||
|
data = response.json()
|
||||||
|
assert len(data["items"]) == 1
|
||||||
|
assert data["items"][0]["title"] == "A"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
def test_list_pagination(
|
||||||
|
authenticated_client: TestClient,
|
||||||
|
db_session: AsyncSession,
|
||||||
|
notification_service: NotificationService,
|
||||||
|
user_a: User,
|
||||||
|
) -> None:
|
||||||
|
async def create_many() -> None:
|
||||||
|
for i in range(25):
|
||||||
|
n = await notification_service.create_notification(
|
||||||
|
db_session,
|
||||||
|
user_a.id,
|
||||||
|
category="instance",
|
||||||
|
severity="info",
|
||||||
|
title=f"Notification {i}",
|
||||||
|
)
|
||||||
|
n.created_at = datetime.now(timezone.utc) - timedelta(seconds=i)
|
||||||
|
await db_session.commit()
|
||||||
|
await db_session.refresh(n)
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
|
||||||
|
asyncio.run(create_many())
|
||||||
|
|
||||||
|
_mint_cookie_for_user(authenticated_client, user_a.id)
|
||||||
|
response = authenticated_client.get("/notifications?limit=10&offset=10")
|
||||||
|
assert response.status_code == 200
|
||||||
|
data = response.json()
|
||||||
|
assert len(data["items"]) == 10
|
||||||
|
assert data["total"] == 25
|
||||||
|
assert data["limit"] == 10
|
||||||
|
assert data["offset"] == 10
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
def test_unread_count_endpoint(
|
||||||
|
authenticated_client: TestClient,
|
||||||
|
db_session: AsyncSession,
|
||||||
|
notification_service: NotificationService,
|
||||||
|
user_a: User,
|
||||||
|
) -> None:
|
||||||
|
async def create_unread() -> None:
|
||||||
|
for _ in range(3):
|
||||||
|
await notification_service.create_notification(
|
||||||
|
db_session,
|
||||||
|
user_a.id,
|
||||||
|
category="instance",
|
||||||
|
severity="info",
|
||||||
|
title="Unread",
|
||||||
|
)
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
|
||||||
|
asyncio.run(create_unread())
|
||||||
|
|
||||||
|
_mint_cookie_for_user(authenticated_client, user_a.id)
|
||||||
|
response = authenticated_client.get("/notifications/unread")
|
||||||
|
assert response.status_code == 200
|
||||||
|
data = response.json()
|
||||||
|
assert data["count"] == 3
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
def test_mark_read_endpoint(
|
||||||
|
authenticated_client: TestClient,
|
||||||
|
db_session: AsyncSession,
|
||||||
|
notification_service: NotificationService,
|
||||||
|
user_a: User,
|
||||||
|
) -> None:
|
||||||
|
async def create_and_get() -> uuid.UUID:
|
||||||
|
n = await notification_service.create_notification(
|
||||||
|
db_session, user_a.id, category="instance", severity="info", title="To read"
|
||||||
|
)
|
||||||
|
return n.id
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
|
||||||
|
nid = asyncio.run(create_and_get())
|
||||||
|
|
||||||
|
_mint_cookie_for_user(authenticated_client, user_a.id)
|
||||||
|
response = authenticated_client.patch(f"/notifications/{nid}/read")
|
||||||
|
assert response.status_code == 200
|
||||||
|
data = response.json()
|
||||||
|
assert data["read_at"] is not None
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
def test_mark_read_404_for_other_user(
|
||||||
|
authenticated_client: TestClient,
|
||||||
|
db_session: AsyncSession,
|
||||||
|
notification_service: NotificationService,
|
||||||
|
user_a: User,
|
||||||
|
user_b: User,
|
||||||
|
) -> None:
|
||||||
|
async def create_and_get() -> uuid.UUID:
|
||||||
|
n = await notification_service.create_notification(
|
||||||
|
db_session,
|
||||||
|
user_a.id,
|
||||||
|
category="instance",
|
||||||
|
severity="info",
|
||||||
|
title="Owned by A",
|
||||||
|
)
|
||||||
|
return n.id
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
|
||||||
|
nid = asyncio.run(create_and_get())
|
||||||
|
|
||||||
|
_mint_cookie_for_user(authenticated_client, user_b.id)
|
||||||
|
response = authenticated_client.patch(f"/notifications/{nid}/read")
|
||||||
|
assert response.status_code == 404
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
def test_mark_all_read_endpoint(
|
||||||
|
authenticated_client: TestClient,
|
||||||
|
db_session: AsyncSession,
|
||||||
|
notification_service: NotificationService,
|
||||||
|
user_a: User,
|
||||||
|
) -> None:
|
||||||
|
async def create_unread() -> None:
|
||||||
|
for _ in range(4):
|
||||||
|
await notification_service.create_notification(
|
||||||
|
db_session,
|
||||||
|
user_a.id,
|
||||||
|
category="instance",
|
||||||
|
severity="info",
|
||||||
|
title="Unread",
|
||||||
|
)
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
|
||||||
|
asyncio.run(create_unread())
|
||||||
|
|
||||||
|
_mint_cookie_for_user(authenticated_client, user_a.id)
|
||||||
|
response = authenticated_client.post("/notifications/mark-all-read")
|
||||||
|
assert response.status_code == 200
|
||||||
|
data = response.json()
|
||||||
|
assert data["marked_count"] == 4
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
def test_dismiss_endpoint(
|
||||||
|
authenticated_client: TestClient,
|
||||||
|
db_session: AsyncSession,
|
||||||
|
notification_service: NotificationService,
|
||||||
|
user_a: User,
|
||||||
|
) -> None:
|
||||||
|
async def create_and_get() -> uuid.UUID:
|
||||||
|
n = await notification_service.create_notification(
|
||||||
|
db_session,
|
||||||
|
user_a.id,
|
||||||
|
category="instance",
|
||||||
|
severity="info",
|
||||||
|
title="To dismiss",
|
||||||
|
)
|
||||||
|
return n.id
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
|
||||||
|
nid = asyncio.run(create_and_get())
|
||||||
|
|
||||||
|
_mint_cookie_for_user(authenticated_client, user_a.id)
|
||||||
|
response = authenticated_client.delete(f"/notifications/{nid}")
|
||||||
|
assert response.status_code == 204
|
||||||
|
|
||||||
|
response = authenticated_client.get("/notifications")
|
||||||
|
data = response.json()
|
||||||
|
assert len(data["items"]) == 0
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
def test_dismiss_404_for_other_user(
|
||||||
|
authenticated_client: TestClient,
|
||||||
|
db_session: AsyncSession,
|
||||||
|
notification_service: NotificationService,
|
||||||
|
user_a: User,
|
||||||
|
user_b: User,
|
||||||
|
) -> None:
|
||||||
|
async def create_and_get() -> uuid.UUID:
|
||||||
|
n = await notification_service.create_notification(
|
||||||
|
db_session,
|
||||||
|
user_a.id,
|
||||||
|
category="instance",
|
||||||
|
severity="info",
|
||||||
|
title="Owned by A",
|
||||||
|
)
|
||||||
|
return n.id
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
|
||||||
|
nid = asyncio.run(create_and_get())
|
||||||
|
|
||||||
|
_mint_cookie_for_user(authenticated_client, user_b.id)
|
||||||
|
response = authenticated_client.delete(f"/notifications/{nid}")
|
||||||
|
assert response.status_code == 404
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.integration
|
||||||
|
def test_mute_categories_filter_in_list(
|
||||||
|
authenticated_client: TestClient,
|
||||||
|
db_session: AsyncSession,
|
||||||
|
notification_service: NotificationService,
|
||||||
|
user_a: User,
|
||||||
|
) -> None:
|
||||||
|
async def setup() -> None:
|
||||||
|
config = UserConfig(
|
||||||
|
user_id=user_a.id, config={"notification_mute_categories": ["instance"]}
|
||||||
|
)
|
||||||
|
db_session.add(config)
|
||||||
|
await db_session.commit()
|
||||||
|
|
||||||
|
await notification_service.create_notification(
|
||||||
|
db_session,
|
||||||
|
user_a.id,
|
||||||
|
category="instance",
|
||||||
|
severity="info",
|
||||||
|
title="Instance",
|
||||||
|
)
|
||||||
|
await notification_service.create_notification(
|
||||||
|
db_session, user_a.id, category="system", severity="info", title="System"
|
||||||
|
)
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
|
||||||
|
asyncio.run(setup())
|
||||||
|
|
||||||
|
_mint_cookie_for_user(authenticated_client, user_a.id)
|
||||||
|
response = authenticated_client.get("/notifications")
|
||||||
|
assert response.status_code == 200
|
||||||
|
data = response.json()
|
||||||
|
assert len(data["items"]) == 1
|
||||||
|
assert data["items"][0]["title"] == "System"
|
||||||
@@ -0,0 +1,395 @@
|
|||||||
|
"""Integration tests for event producer → notification creation flow."""
|
||||||
|
|
||||||
|
import uuid
|
||||||
|
from collections.abc import Generator
|
||||||
|
from unittest.mock import patch
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
import pytest_asyncio
|
||||||
|
from sqlalchemy import select
|
||||||
|
from sqlalchemy.ext.asyncio import AsyncSession
|
||||||
|
|
||||||
|
from src.models.git_repository import GitRepository
|
||||||
|
from src.models.notification import Notification
|
||||||
|
from src.models.project import Project
|
||||||
|
from src.models.tool_instance import ToolInstance
|
||||||
|
from src.models.tool_type import ToolType
|
||||||
|
from src.models.user import User
|
||||||
|
from src.services.event_bus import InstanceEventBus, InstanceEventPayload
|
||||||
|
from src.services.health_monitor import HealthSnapshot
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def event_bus() -> Generator[InstanceEventBus, None, None]:
|
||||||
|
"""Provide a fresh EventBus instance."""
|
||||||
|
bus = InstanceEventBus()
|
||||||
|
bus._reset_for_testing()
|
||||||
|
yield bus
|
||||||
|
bus._reset_for_testing()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest_asyncio.fixture
|
||||||
|
async def test_instance(db_session: AsyncSession) -> ToolInstance:
|
||||||
|
"""Create a complete tool instance with all required relations."""
|
||||||
|
user = User(
|
||||||
|
id=uuid.uuid4(),
|
||||||
|
email="owner@headquarter.local",
|
||||||
|
name="Owner",
|
||||||
|
authentik_id=f"authentik-{uuid.uuid4()}",
|
||||||
|
avatar_url=None,
|
||||||
|
)
|
||||||
|
db_session.add(user)
|
||||||
|
await db_session.commit()
|
||||||
|
|
||||||
|
project = Project(
|
||||||
|
id=uuid.uuid4(),
|
||||||
|
name="test-project",
|
||||||
|
description="Test",
|
||||||
|
owner_id=user.id,
|
||||||
|
)
|
||||||
|
repo = GitRepository(
|
||||||
|
id=uuid.uuid4(),
|
||||||
|
name="test-repo",
|
||||||
|
path="/tmp/test-repo",
|
||||||
|
project_id=project.id,
|
||||||
|
owner_id=user.id,
|
||||||
|
remote_url="https://github.com/test/repo.git",
|
||||||
|
)
|
||||||
|
tool_type = ToolType(
|
||||||
|
id=uuid.uuid4(),
|
||||||
|
name="test-tool",
|
||||||
|
display_name="Test Tool",
|
||||||
|
category="other",
|
||||||
|
interface_type="web",
|
||||||
|
requires_port=True,
|
||||||
|
default_port=8080,
|
||||||
|
definition_type="legacy",
|
||||||
|
compose_template="version: '3.8'\nservices:\n app:\n image: alpine\n command: sleep 3600\n",
|
||||||
|
)
|
||||||
|
db_session.add_all([project, repo, tool_type])
|
||||||
|
await db_session.commit()
|
||||||
|
|
||||||
|
instance = ToolInstance(
|
||||||
|
id=uuid.uuid4(),
|
||||||
|
name="test-instance",
|
||||||
|
display_name="Test Instance",
|
||||||
|
tool_type_id=tool_type.id,
|
||||||
|
repository_id=repo.id,
|
||||||
|
project_id=project.id,
|
||||||
|
owner_id=user.id,
|
||||||
|
status="running",
|
||||||
|
compose_path="/tmp/test-compose.yml",
|
||||||
|
port=8080,
|
||||||
|
)
|
||||||
|
db_session.add(instance)
|
||||||
|
await db_session.commit()
|
||||||
|
return instance
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.integration
|
||||||
|
async def test_lifecycle_started_intermediate_skips_notification(
|
||||||
|
db_session: AsyncSession,
|
||||||
|
event_bus: InstanceEventBus,
|
||||||
|
test_instance: ToolInstance,
|
||||||
|
) -> None:
|
||||||
|
"""Intermediate 'starting' state does NOT create a notification."""
|
||||||
|
received: list[InstanceEventPayload] = []
|
||||||
|
|
||||||
|
def subscriber(payload: InstanceEventPayload) -> None:
|
||||||
|
received.append(payload)
|
||||||
|
|
||||||
|
event_bus.subscribe("instance.started", subscriber)
|
||||||
|
|
||||||
|
from src.services.lifecycle_hooks import publish_lifecycle_event
|
||||||
|
|
||||||
|
await publish_lifecycle_event(
|
||||||
|
event_bus=event_bus,
|
||||||
|
session=db_session,
|
||||||
|
instance=test_instance,
|
||||||
|
event_type="instance.started",
|
||||||
|
status="starting",
|
||||||
|
message="Container starting...",
|
||||||
|
)
|
||||||
|
|
||||||
|
# Event still published
|
||||||
|
assert len(received) == 1
|
||||||
|
|
||||||
|
# No notification created for intermediate state
|
||||||
|
result = await db_session.execute(
|
||||||
|
select(Notification).where(Notification.user_id == test_instance.owner_id)
|
||||||
|
)
|
||||||
|
notifications = list(result.scalars().all())
|
||||||
|
assert len(notifications) == 0
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.integration
|
||||||
|
async def test_lifecycle_running_creates_notification(
|
||||||
|
db_session: AsyncSession,
|
||||||
|
event_bus: InstanceEventBus,
|
||||||
|
test_instance: ToolInstance,
|
||||||
|
) -> None:
|
||||||
|
"""Successful terminal state (running) creates a notification."""
|
||||||
|
from src.services.lifecycle_hooks import publish_lifecycle_event
|
||||||
|
|
||||||
|
await publish_lifecycle_event(
|
||||||
|
event_bus=event_bus,
|
||||||
|
session=db_session,
|
||||||
|
instance=test_instance,
|
||||||
|
event_type="instance.health_changed",
|
||||||
|
status="running",
|
||||||
|
message="Container running",
|
||||||
|
)
|
||||||
|
|
||||||
|
result = await db_session.execute(
|
||||||
|
select(Notification).where(Notification.user_id == test_instance.owner_id)
|
||||||
|
)
|
||||||
|
notifications = list(result.scalars().all())
|
||||||
|
assert len(notifications) == 1
|
||||||
|
n = notifications[0]
|
||||||
|
assert n.category == "instance"
|
||||||
|
assert n.severity == "success"
|
||||||
|
assert n.title == "Container ready"
|
||||||
|
assert n.source_type == "tool_instances"
|
||||||
|
assert n.source_id == test_instance.id
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.integration
|
||||||
|
async def test_health_monitor_error_creates_notification(
|
||||||
|
db_session: AsyncSession,
|
||||||
|
event_bus: InstanceEventBus,
|
||||||
|
test_instance: ToolInstance,
|
||||||
|
) -> None:
|
||||||
|
"""Simulating a health monitor crash creates an error notification."""
|
||||||
|
from src.services.health_monitor import HealthMonitor
|
||||||
|
|
||||||
|
monitor = HealthMonitor(event_bus)
|
||||||
|
|
||||||
|
received: list[InstanceEventPayload] = []
|
||||||
|
|
||||||
|
def subscriber(payload: InstanceEventPayload) -> None:
|
||||||
|
received.append(payload)
|
||||||
|
|
||||||
|
event_bus.subscribe("instance.error", subscriber)
|
||||||
|
|
||||||
|
with patch(
|
||||||
|
"src.services.health_monitor.get_container_status",
|
||||||
|
return_value={"status": "exited", "exit_code": 137, "health": None},
|
||||||
|
):
|
||||||
|
await monitor._check_instance(db_session, test_instance)
|
||||||
|
|
||||||
|
# Event published
|
||||||
|
assert len(received) == 1
|
||||||
|
|
||||||
|
# Notification created
|
||||||
|
result = await db_session.execute(
|
||||||
|
select(Notification).where(Notification.user_id == test_instance.owner_id)
|
||||||
|
)
|
||||||
|
notifications = list(result.scalars().all())
|
||||||
|
assert len(notifications) == 1
|
||||||
|
n = notifications[0]
|
||||||
|
assert n.category == "instance"
|
||||||
|
assert n.severity == "error"
|
||||||
|
assert n.source_type == "tool_instances"
|
||||||
|
assert n.source_id == test_instance.id
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.integration
|
||||||
|
async def test_notification_failure_does_not_block_event_pipeline(
|
||||||
|
db_session: AsyncSession,
|
||||||
|
event_bus: InstanceEventBus,
|
||||||
|
test_instance: ToolInstance,
|
||||||
|
) -> None:
|
||||||
|
"""If NotificationService raises, the event is still published and no exception escapes."""
|
||||||
|
received: list[InstanceEventPayload] = []
|
||||||
|
|
||||||
|
def subscriber(payload: InstanceEventPayload) -> None:
|
||||||
|
received.append(payload)
|
||||||
|
|
||||||
|
event_bus.subscribe("instance.started", subscriber)
|
||||||
|
|
||||||
|
from src.services.lifecycle_hooks import publish_lifecycle_event
|
||||||
|
|
||||||
|
with patch(
|
||||||
|
"src.services.lifecycle_hooks.notification_service.create_notification",
|
||||||
|
side_effect=RuntimeError("DB is down"),
|
||||||
|
):
|
||||||
|
# Should not raise
|
||||||
|
await publish_lifecycle_event(
|
||||||
|
event_bus=event_bus,
|
||||||
|
session=db_session,
|
||||||
|
instance=test_instance,
|
||||||
|
event_type="instance.started",
|
||||||
|
status="starting",
|
||||||
|
message="Container started",
|
||||||
|
)
|
||||||
|
|
||||||
|
assert len(received) == 1
|
||||||
|
assert received[0]["event"] == "instance.started"
|
||||||
|
|
||||||
|
# No notification should have been created
|
||||||
|
result = await db_session.execute(
|
||||||
|
select(Notification).where(Notification.user_id == test_instance.owner_id)
|
||||||
|
)
|
||||||
|
assert result.scalar_one_or_none() is None
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.integration
|
||||||
|
async def test_notification_ownership_matches_instance_owner(
|
||||||
|
db_session: AsyncSession,
|
||||||
|
event_bus: InstanceEventBus,
|
||||||
|
) -> None:
|
||||||
|
"""Notification user_id matches the instance owner, not any caller."""
|
||||||
|
# Create a caller user (simulates the user making an API request)
|
||||||
|
caller = User(
|
||||||
|
id=uuid.uuid4(),
|
||||||
|
email="caller@headquarter.local",
|
||||||
|
name="Caller",
|
||||||
|
authentik_id=f"authentik-{uuid.uuid4()}",
|
||||||
|
avatar_url=None,
|
||||||
|
)
|
||||||
|
db_session.add(caller)
|
||||||
|
await db_session.commit()
|
||||||
|
|
||||||
|
# Create the actual owner
|
||||||
|
owner = User(
|
||||||
|
id=uuid.uuid4(),
|
||||||
|
email="owner@headquarter.local",
|
||||||
|
name="Owner",
|
||||||
|
authentik_id=f"authentik-{uuid.uuid4()}",
|
||||||
|
avatar_url=None,
|
||||||
|
)
|
||||||
|
db_session.add(owner)
|
||||||
|
await db_session.commit()
|
||||||
|
|
||||||
|
project = Project(
|
||||||
|
id=uuid.uuid4(),
|
||||||
|
name="test-project",
|
||||||
|
description="Test",
|
||||||
|
owner_id=owner.id,
|
||||||
|
)
|
||||||
|
repo = GitRepository(
|
||||||
|
id=uuid.uuid4(),
|
||||||
|
name="test-repo",
|
||||||
|
path="/tmp/test-repo",
|
||||||
|
project_id=project.id,
|
||||||
|
owner_id=owner.id,
|
||||||
|
remote_url="https://github.com/test/repo.git",
|
||||||
|
)
|
||||||
|
tool_type = ToolType(
|
||||||
|
id=uuid.uuid4(),
|
||||||
|
name="test-tool",
|
||||||
|
display_name="Test Tool",
|
||||||
|
category="other",
|
||||||
|
interface_type="web",
|
||||||
|
requires_port=True,
|
||||||
|
default_port=8080,
|
||||||
|
definition_type="legacy",
|
||||||
|
compose_template="version: '3.8'\nservices:\n app:\n image: alpine\n command: sleep 3600\n",
|
||||||
|
)
|
||||||
|
db_session.add_all([project, repo, tool_type])
|
||||||
|
await db_session.commit()
|
||||||
|
|
||||||
|
instance = ToolInstance(
|
||||||
|
id=uuid.uuid4(),
|
||||||
|
name="test-instance",
|
||||||
|
display_name="Test Instance",
|
||||||
|
tool_type_id=tool_type.id,
|
||||||
|
repository_id=repo.id,
|
||||||
|
project_id=project.id,
|
||||||
|
owner_id=owner.id,
|
||||||
|
status="running",
|
||||||
|
compose_path="/tmp/test-compose.yml",
|
||||||
|
port=8080,
|
||||||
|
)
|
||||||
|
db_session.add(instance)
|
||||||
|
await db_session.commit()
|
||||||
|
|
||||||
|
from src.services.lifecycle_hooks import publish_lifecycle_event
|
||||||
|
|
||||||
|
await publish_lifecycle_event(
|
||||||
|
event_bus=event_bus,
|
||||||
|
session=db_session,
|
||||||
|
instance=instance,
|
||||||
|
event_type="instance.health_changed",
|
||||||
|
status="running",
|
||||||
|
message="Container running",
|
||||||
|
)
|
||||||
|
|
||||||
|
result = await db_session.execute(
|
||||||
|
select(Notification).where(Notification.source_id == instance.id)
|
||||||
|
)
|
||||||
|
n = result.scalar_one()
|
||||||
|
assert n.user_id == owner.id
|
||||||
|
assert n.user_id != caller.id
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.integration
|
||||||
|
async def test_lifecycle_error_creates_error_notification(
|
||||||
|
db_session: AsyncSession,
|
||||||
|
event_bus: InstanceEventBus,
|
||||||
|
test_instance: ToolInstance,
|
||||||
|
) -> None:
|
||||||
|
"""An instance.error lifecycle event creates a severity=error notification."""
|
||||||
|
from src.services.lifecycle_hooks import publish_lifecycle_event
|
||||||
|
|
||||||
|
await publish_lifecycle_event(
|
||||||
|
event_bus=event_bus,
|
||||||
|
session=db_session,
|
||||||
|
instance=test_instance,
|
||||||
|
event_type="instance.error",
|
||||||
|
status="error",
|
||||||
|
message="Container failed",
|
||||||
|
)
|
||||||
|
|
||||||
|
result = await db_session.execute(
|
||||||
|
select(Notification).where(Notification.user_id == test_instance.owner_id)
|
||||||
|
)
|
||||||
|
n = result.scalar_one()
|
||||||
|
assert n.severity == "error"
|
||||||
|
assert n.title == "Container error"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.integration
|
||||||
|
async def test_health_monitor_unhealthy_creates_warning_notification(
|
||||||
|
db_session: AsyncSession,
|
||||||
|
event_bus: InstanceEventBus,
|
||||||
|
test_instance: ToolInstance,
|
||||||
|
) -> None:
|
||||||
|
"""Health monitor marking instance unhealthy creates severity=warning notification."""
|
||||||
|
from src.services.health_monitor import HealthMonitor
|
||||||
|
|
||||||
|
monitor = HealthMonitor(event_bus)
|
||||||
|
monitor._last_known_state[test_instance.id] = HealthSnapshot(
|
||||||
|
container_status="running",
|
||||||
|
container_healthy=None,
|
||||||
|
tunnel_healthy=True,
|
||||||
|
exit_code=None,
|
||||||
|
)
|
||||||
|
test_instance.public_url = "https://example.trycloudflare.com"
|
||||||
|
|
||||||
|
with (
|
||||||
|
patch(
|
||||||
|
"src.services.health_monitor.get_container_status",
|
||||||
|
return_value={"status": "running", "exit_code": None, "health": "healthy"},
|
||||||
|
),
|
||||||
|
patch(
|
||||||
|
"src.services.health_monitor.check_tunnel_health",
|
||||||
|
return_value={"healthy": False, "tunnel_status": "error_response"},
|
||||||
|
),
|
||||||
|
):
|
||||||
|
await monitor._check_instance(db_session, test_instance)
|
||||||
|
|
||||||
|
result = await db_session.execute(
|
||||||
|
select(Notification).where(Notification.user_id == test_instance.owner_id)
|
||||||
|
)
|
||||||
|
n = result.scalar_one()
|
||||||
|
assert n.category == "health"
|
||||||
|
assert n.severity == "warning"
|
||||||
|
assert n.title == "Container unhealthy"
|
||||||
@@ -3,6 +3,7 @@ from datetime import UTC, datetime, timedelta
|
|||||||
import asyncio
|
import asyncio
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
from fastapi.testclient import TestClient
|
||||||
from sqlalchemy import text
|
from sqlalchemy import text
|
||||||
from sqlalchemy.ext.asyncio import create_async_engine, async_sessionmaker
|
from sqlalchemy.ext.asyncio import create_async_engine, async_sessionmaker
|
||||||
|
|
||||||
|
|||||||
@@ -1,256 +0,0 @@
|
|||||||
import uuid
|
|
||||||
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
|
|
||||||
@@ -1,4 +1,3 @@
|
|||||||
import uuid
|
|
||||||
import pytest
|
import pytest
|
||||||
from fastapi.testclient import TestClient
|
from fastapi.testclient import TestClient
|
||||||
|
|
||||||
@@ -7,7 +6,9 @@ from fastapi.testclient import TestClient
|
|||||||
class TestToolTypesAPIExtended:
|
class TestToolTypesAPIExtended:
|
||||||
"""Integration tests for tool types API with new fields."""
|
"""Integration tests for tool types API with new fields."""
|
||||||
|
|
||||||
def test_create_tool_type_with_dockerfile(self, authenticated_client: TestClient) -> None:
|
def test_create_tool_type_with_dockerfile(
|
||||||
|
self, authenticated_client: TestClient
|
||||||
|
) -> None:
|
||||||
"""Test creating a tool type with dockerfile definition."""
|
"""Test creating a tool type with dockerfile definition."""
|
||||||
response = authenticated_client.post(
|
response = authenticated_client.post(
|
||||||
"/tool-types",
|
"/tool-types",
|
||||||
@@ -28,7 +29,9 @@ class TestToolTypesAPIExtended:
|
|||||||
assert data["definition_type"] == "dockerfile"
|
assert data["definition_type"] == "dockerfile"
|
||||||
assert data["dockerfile_template"] == "FROM python:3.11\nRUN pip install flask"
|
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:
|
def test_create_tool_type_with_readiness_probe(
|
||||||
|
self, authenticated_client: TestClient
|
||||||
|
) -> None:
|
||||||
"""Test creating a tool type with readiness probe."""
|
"""Test creating a tool type with readiness probe."""
|
||||||
response = authenticated_client.post(
|
response = authenticated_client.post(
|
||||||
"/tool-types",
|
"/tool-types",
|
||||||
@@ -53,7 +56,9 @@ class TestToolTypesAPIExtended:
|
|||||||
assert data["readiness_probe"]["command"] == "curl -f http://localhost:8080"
|
assert data["readiness_probe"]["command"] == "curl -f http://localhost:8080"
|
||||||
assert data["readiness_probe"]["timeout"] == 30
|
assert data["readiness_probe"]["timeout"] == 30
|
||||||
|
|
||||||
def test_create_tool_type_invalid_definition_type(self, authenticated_client: TestClient) -> None:
|
def test_create_tool_type_invalid_definition_type(
|
||||||
|
self, authenticated_client: TestClient
|
||||||
|
) -> None:
|
||||||
"""Test that invalid definition types are rejected."""
|
"""Test that invalid definition types are rejected."""
|
||||||
response = authenticated_client.post(
|
response = authenticated_client.post(
|
||||||
"/tool-types",
|
"/tool-types",
|
||||||
@@ -68,7 +73,9 @@ class TestToolTypesAPIExtended:
|
|||||||
)
|
)
|
||||||
assert response.status_code == 422
|
assert response.status_code == 422
|
||||||
|
|
||||||
def test_create_tool_type_dockerfile_without_template(self, authenticated_client: TestClient) -> None:
|
def test_create_tool_type_dockerfile_without_template(
|
||||||
|
self, authenticated_client: TestClient
|
||||||
|
) -> None:
|
||||||
"""Test that dockerfile type requires dockerfile_template."""
|
"""Test that dockerfile type requires dockerfile_template."""
|
||||||
response = authenticated_client.post(
|
response = authenticated_client.post(
|
||||||
"/tool-types",
|
"/tool-types",
|
||||||
@@ -82,7 +89,9 @@ class TestToolTypesAPIExtended:
|
|||||||
)
|
)
|
||||||
assert response.status_code == 422
|
assert response.status_code == 422
|
||||||
|
|
||||||
def test_update_tool_type_with_new_fields(self, authenticated_client: TestClient) -> None:
|
def test_update_tool_type_with_new_fields(
|
||||||
|
self, authenticated_client: TestClient
|
||||||
|
) -> None:
|
||||||
"""Test updating a tool type with new fields."""
|
"""Test updating a tool type with new fields."""
|
||||||
# Create tool type first
|
# Create tool type first
|
||||||
create_response = authenticated_client.post(
|
create_response = authenticated_client.post(
|
||||||
@@ -113,7 +122,9 @@ class TestToolTypesAPIExtended:
|
|||||||
assert response.status_code == 200
|
assert response.status_code == 200
|
||||||
data = response.json()
|
data = response.json()
|
||||||
assert data["display_name"] == "Updated Name"
|
assert data["display_name"] == "Updated Name"
|
||||||
assert data["readiness_probe"]["command"] == "curl -f http://localhost:8080/health"
|
assert (
|
||||||
|
data["readiness_probe"]["command"] == "curl -f http://localhost:8080/health"
|
||||||
|
)
|
||||||
|
|
||||||
def test_validate_tool_type_compose(self, authenticated_client: TestClient) -> None:
|
def test_validate_tool_type_compose(self, authenticated_client: TestClient) -> None:
|
||||||
"""Test validating compose template."""
|
"""Test validating compose template."""
|
||||||
@@ -128,7 +139,9 @@ class TestToolTypesAPIExtended:
|
|||||||
data = response.json()
|
data = response.json()
|
||||||
assert data["valid"] is True
|
assert data["valid"] is True
|
||||||
|
|
||||||
def test_validate_tool_type_invalid_compose(self, authenticated_client: TestClient) -> None:
|
def test_validate_tool_type_invalid_compose(
|
||||||
|
self, authenticated_client: TestClient
|
||||||
|
) -> None:
|
||||||
"""Test validating invalid compose template."""
|
"""Test validating invalid compose template."""
|
||||||
response = authenticated_client.post(
|
response = authenticated_client.post(
|
||||||
"/tool-types/validate",
|
"/tool-types/validate",
|
||||||
@@ -142,7 +155,9 @@ class TestToolTypesAPIExtended:
|
|||||||
assert data["valid"] is False
|
assert data["valid"] is False
|
||||||
assert "errors" in data
|
assert "errors" in data
|
||||||
|
|
||||||
def test_validate_tool_type_dockerfile(self, authenticated_client: TestClient) -> None:
|
def test_validate_tool_type_dockerfile(
|
||||||
|
self, authenticated_client: TestClient
|
||||||
|
) -> None:
|
||||||
"""Test validating dockerfile template."""
|
"""Test validating dockerfile template."""
|
||||||
response = authenticated_client.post(
|
response = authenticated_client.post(
|
||||||
"/tool-types/validate",
|
"/tool-types/validate",
|
||||||
@@ -155,7 +170,9 @@ class TestToolTypesAPIExtended:
|
|||||||
data = response.json()
|
data = response.json()
|
||||||
assert data["valid"] is True
|
assert data["valid"] is True
|
||||||
|
|
||||||
def test_get_tool_type_returns_new_fields(self, authenticated_client: TestClient) -> None:
|
def test_get_tool_type_returns_new_fields(
|
||||||
|
self, authenticated_client: TestClient
|
||||||
|
) -> None:
|
||||||
"""Test that GET returns new fields."""
|
"""Test that GET returns new fields."""
|
||||||
# Create tool type with all fields
|
# Create tool type with all fields
|
||||||
create_response = authenticated_client.post(
|
create_response = authenticated_client.post(
|
||||||
@@ -167,7 +184,7 @@ class TestToolTypesAPIExtended:
|
|||||||
"interfaces": ["web", "terminal"],
|
"interfaces": ["web", "terminal"],
|
||||||
"default_port": 8443,
|
"default_port": 8443,
|
||||||
"definition_type": "compose",
|
"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\"",
|
"compose_template": "version: '3.8'\nservices:\n app:\n image: code-server\n command: --bind-addr 0.0.0.0:8443\n ports:\n - '8443:8443'\n volumes:\n - \"{{REPO_PATH}}:/workspace\"",
|
||||||
"readiness_probe": {
|
"readiness_probe": {
|
||||||
"command": "curl -f http://localhost:8443",
|
"command": "curl -f http://localhost:8443",
|
||||||
"timeout": 30,
|
"timeout": 30,
|
||||||
@@ -187,7 +204,9 @@ class TestToolTypesAPIExtended:
|
|||||||
assert data["interfaces"] == ["web", "terminal"]
|
assert data["interfaces"] == ["web", "terminal"]
|
||||||
assert "readiness_probe" in data
|
assert "readiness_probe" in data
|
||||||
|
|
||||||
def test_create_tool_type_without_port_fails(self, authenticated_client: TestClient) -> None:
|
def test_create_tool_type_without_port_fails(
|
||||||
|
self, authenticated_client: TestClient
|
||||||
|
) -> None:
|
||||||
"""Test that creating a tool type without default_port fails validation."""
|
"""Test that creating a tool type without default_port fails validation."""
|
||||||
response = authenticated_client.post(
|
response = authenticated_client.post(
|
||||||
"/tool-types",
|
"/tool-types",
|
||||||
@@ -205,7 +224,9 @@ class TestToolTypesAPIExtended:
|
|||||||
data = response.json()
|
data = response.json()
|
||||||
assert "default_port" in str(data)
|
assert "default_port" in str(data)
|
||||||
|
|
||||||
def test_create_tool_type_with_port_mismatch_fails(self, authenticated_client: TestClient) -> None:
|
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."""
|
"""Test that port mismatch between default_port and compose template fails."""
|
||||||
response = authenticated_client.post(
|
response = authenticated_client.post(
|
||||||
"/tool-types",
|
"/tool-types",
|
||||||
@@ -221,5 +242,85 @@ class TestToolTypesAPIExtended:
|
|||||||
},
|
},
|
||||||
)
|
)
|
||||||
assert response.status_code == 422
|
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()
|
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)
|
assert "Port 9999 is not exposed" in str(data)
|
||||||
|
|||||||
@@ -0,0 +1,203 @@
|
|||||||
|
"""Unit tests for TerminalManager multi-session support."""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import uuid
|
||||||
|
from unittest.mock import AsyncMock, patch
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from src.services.terminal_manager import MaxSessionsExceededError, TerminalManager
|
||||||
|
from src.services.terminal_session import TerminalSession
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def manager() -> TerminalManager:
|
||||||
|
"""Provide a fresh TerminalManager instance for each test."""
|
||||||
|
tm = TerminalManager()
|
||||||
|
# Cancel the background idle check to avoid side effects
|
||||||
|
if tm._idle_check_task and not tm._idle_check_task.done():
|
||||||
|
tm._idle_check_task.cancel()
|
||||||
|
return tm
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def mock_terminal_session(monkeypatch) -> None:
|
||||||
|
"""Monkeypatch TerminalSession.start and is_alive for unit tests."""
|
||||||
|
|
||||||
|
async def fake_start(self, startup_command=None):
|
||||||
|
self.last_activity = __import__("time").time()
|
||||||
|
|
||||||
|
monkeypatch.setattr(TerminalSession, "start", fake_start)
|
||||||
|
monkeypatch.setattr(TerminalSession, "is_alive", lambda self: True)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def instance_id() -> uuid.UUID:
|
||||||
|
return uuid.uuid4()
|
||||||
|
|
||||||
|
|
||||||
|
class FakeWebSocket:
|
||||||
|
"""Minimal fake WebSocket for testing attach/detach behavior."""
|
||||||
|
|
||||||
|
def __init__(self, name: str = "ws") -> None:
|
||||||
|
self.name = name
|
||||||
|
self.closed = False
|
||||||
|
self.close_code: int | None = None
|
||||||
|
self.close_reason: str | None = None
|
||||||
|
self._sent: list[bytes] = []
|
||||||
|
|
||||||
|
async def close(self, code: int = 1000, reason: str = "") -> None:
|
||||||
|
self.closed = True
|
||||||
|
self.close_code = code
|
||||||
|
self.close_reason = reason
|
||||||
|
|
||||||
|
async def send_bytes(self, data: bytes) -> None:
|
||||||
|
self._sent.append(data)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_create_session_increases_count(
|
||||||
|
manager: TerminalManager,
|
||||||
|
mock_terminal_session,
|
||||||
|
instance_id: uuid.UUID,
|
||||||
|
) -> None:
|
||||||
|
"""Creating sessions increments the per-instance count."""
|
||||||
|
assert len(manager.get_sessions_for_instance(str(instance_id))) == 0
|
||||||
|
|
||||||
|
session1 = await manager.create_session(instance_id, "container-1")
|
||||||
|
assert len(manager.get_sessions_for_instance(str(instance_id))) == 1
|
||||||
|
assert session1.session_id in [
|
||||||
|
s.session_id for s in manager.get_sessions_for_instance(str(instance_id))
|
||||||
|
]
|
||||||
|
|
||||||
|
session2 = await manager.create_session(instance_id, "container-1")
|
||||||
|
assert len(manager.get_sessions_for_instance(str(instance_id))) == 2
|
||||||
|
|
||||||
|
# Verify sessions are distinct
|
||||||
|
assert session1.session_id != session2.session_id
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_create_session_enforces_max_5(
|
||||||
|
manager: TerminalManager,
|
||||||
|
mock_terminal_session,
|
||||||
|
instance_id: uuid.UUID,
|
||||||
|
) -> None:
|
||||||
|
"""The 6th session creation raises MaxSessionsExceededError."""
|
||||||
|
for i in range(5):
|
||||||
|
await manager.create_session(instance_id, f"container-{i}")
|
||||||
|
|
||||||
|
assert len(manager.get_sessions_for_instance(str(instance_id))) == 5
|
||||||
|
|
||||||
|
with pytest.raises(MaxSessionsExceededError):
|
||||||
|
await manager.create_session(instance_id, "container-overflow")
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_get_sessions_for_instance_filters_by_instance(
|
||||||
|
manager: TerminalManager,
|
||||||
|
mock_terminal_session,
|
||||||
|
) -> None:
|
||||||
|
"""get_sessions_for_instance returns only sessions for the requested instance."""
|
||||||
|
instance_a = uuid.uuid4()
|
||||||
|
instance_b = uuid.uuid4()
|
||||||
|
|
||||||
|
await manager.create_session(instance_a, "container-a")
|
||||||
|
await manager.create_session(instance_a, "container-a2")
|
||||||
|
await manager.create_session(instance_b, "container-b")
|
||||||
|
|
||||||
|
assert len(manager.get_sessions_for_instance(str(instance_a))) == 2
|
||||||
|
assert len(manager.get_sessions_for_instance(str(instance_b))) == 1
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_close_session_removes_from_dict(
|
||||||
|
manager: TerminalManager,
|
||||||
|
mock_terminal_session,
|
||||||
|
instance_id: uuid.UUID,
|
||||||
|
) -> None:
|
||||||
|
"""close_session removes the key from _sessions and marks DB closed."""
|
||||||
|
session = await manager.create_session(instance_id, "container-1")
|
||||||
|
session_id = session.session_id
|
||||||
|
|
||||||
|
assert manager.get_session(str(instance_id), session_id) is not None
|
||||||
|
|
||||||
|
with patch.object(manager, "_mark_closed_in_db", new=AsyncMock()) as mock_mark:
|
||||||
|
await manager.close_session(str(instance_id), session_id)
|
||||||
|
# Give the fire-and-forget task a chance to be scheduled
|
||||||
|
await asyncio.sleep(0)
|
||||||
|
|
||||||
|
assert manager.get_session(str(instance_id), session_id) is None
|
||||||
|
mock_mark.assert_called_once_with(session_id)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_attach_websocket_only_closes_same_session(
|
||||||
|
manager: TerminalManager,
|
||||||
|
mock_terminal_session,
|
||||||
|
instance_id: uuid.UUID,
|
||||||
|
) -> None:
|
||||||
|
"""Attaching to session A must not close WebSockets on session B."""
|
||||||
|
session_a = await manager.create_session(instance_id, "container-1")
|
||||||
|
session_b = await manager.create_session(instance_id, "container-1")
|
||||||
|
|
||||||
|
ws_a1 = FakeWebSocket("ws-a1")
|
||||||
|
ws_b1 = FakeWebSocket("ws-b1")
|
||||||
|
|
||||||
|
# Manually attach websockets (simulate prior connections)
|
||||||
|
session_a.attach_websocket(ws_a1)
|
||||||
|
session_b.attach_websocket(ws_b1)
|
||||||
|
|
||||||
|
# Now attach a new websocket to session_a
|
||||||
|
ws_a2 = FakeWebSocket("ws-a2")
|
||||||
|
await manager.attach_websocket(session_a, ws_a2)
|
||||||
|
|
||||||
|
# ws_a1 should have been closed because it's on the same session
|
||||||
|
assert ws_a1.closed is True
|
||||||
|
|
||||||
|
# ws_b1 should NOT have been closed because it's on a different session
|
||||||
|
assert ws_b1.closed is False
|
||||||
|
|
||||||
|
# ws_a2 should be attached and receive buffer
|
||||||
|
assert ws_a2 in session_a._websockets
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_default_session_keyed_separately(
|
||||||
|
manager: TerminalManager,
|
||||||
|
mock_terminal_session,
|
||||||
|
instance_id: uuid.UUID,
|
||||||
|
) -> None:
|
||||||
|
"""Default session uses 'default' session_id and does not collide with named sessions."""
|
||||||
|
default_session = await manager.get_or_create_session(instance_id, "container-1")
|
||||||
|
explicit_session = await manager.create_session(instance_id, "container-1")
|
||||||
|
|
||||||
|
# Both should exist
|
||||||
|
assert manager.get_session(str(instance_id), "default") is default_session
|
||||||
|
assert (
|
||||||
|
manager.get_session(str(instance_id), explicit_session.session_id)
|
||||||
|
is explicit_session
|
||||||
|
)
|
||||||
|
|
||||||
|
# They should be different objects
|
||||||
|
assert default_session.session_id != explicit_session.session_id
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_idle_cleanup_updates_db_status(
|
||||||
|
manager: TerminalManager,
|
||||||
|
mock_terminal_session,
|
||||||
|
instance_id: uuid.UUID,
|
||||||
|
) -> None:
|
||||||
|
"""Idle cleanup removes sessions from dict and calls DB update."""
|
||||||
|
session = await manager.create_session(instance_id, "container-1")
|
||||||
|
session_id = session.session_id
|
||||||
|
|
||||||
|
# Make session appear idle (no websockets, old last_activity)
|
||||||
|
session.last_activity = 0
|
||||||
|
|
||||||
|
with patch.object(manager, "_mark_closed_in_db", new=AsyncMock()) as mock_mark:
|
||||||
|
await manager._cleanup_idle_sessions()
|
||||||
|
|
||||||
|
assert manager.get_session(str(instance_id), session_id) is None
|
||||||
|
mock_mark.assert_called_once_with(session_id)
|
||||||
@@ -6,13 +6,16 @@ from src.models.config_profile import ConfigProfile, ConfigProfileInclude
|
|||||||
from src.services.config_profile_resolver import (
|
from src.services.config_profile_resolver import (
|
||||||
ConfigProfileCycleError,
|
ConfigProfileCycleError,
|
||||||
ConfigProfileNotFoundError,
|
ConfigProfileNotFoundError,
|
||||||
|
ResolvedMount,
|
||||||
ResolvedProfile,
|
ResolvedProfile,
|
||||||
|
apply_resolved_profile,
|
||||||
check_include_cycle,
|
check_include_cycle,
|
||||||
resolve_profile,
|
resolve_profile,
|
||||||
_merge_env_vars,
|
_merge_env_vars,
|
||||||
_merge_files,
|
_merge_files,
|
||||||
_merge_mounts,
|
_merge_mounts,
|
||||||
_merge_runtime_hints,
|
_merge_runtime_hints,
|
||||||
|
_merge_git_mounts,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@@ -62,7 +65,6 @@ class TestMergeFunctions:
|
|||||||
|
|
||||||
def test_merge_mounts_basic(self) -> None:
|
def test_merge_mounts_basic(self) -> None:
|
||||||
"""Test basic mount merging."""
|
"""Test basic mount merging."""
|
||||||
from src.services.config_profile_resolver import ResolvedMount
|
|
||||||
result = _merge_mounts(
|
result = _merge_mounts(
|
||||||
{},
|
{},
|
||||||
[{"target": "/app", "mode": "rw", "files": {"a.txt": "content"}}],
|
[{"target": "/app", "mode": "rw", "files": {"a.txt": "content"}}],
|
||||||
@@ -76,6 +78,7 @@ class TestMergeFunctions:
|
|||||||
def test_merge_mounts_file_override(self) -> None:
|
def test_merge_mounts_file_override(self) -> None:
|
||||||
"""Test mount file map merging with overrides."""
|
"""Test mount file map merging with overrides."""
|
||||||
from src.services.config_profile_resolver import ResolvedMount
|
from src.services.config_profile_resolver import ResolvedMount
|
||||||
|
|
||||||
result = _merge_mounts(
|
result = _merge_mounts(
|
||||||
{"/app": ResolvedMount(target="/app", mode="rw", files={"a.txt": "old"})},
|
{"/app": ResolvedMount(target="/app", mode="rw", files={"a.txt": "old"})},
|
||||||
[{"target": "/app", "mode": "rw", "files": {"a.txt": "new"}}],
|
[{"target": "/app", "mode": "rw", "files": {"a.txt": "new"}}],
|
||||||
@@ -87,6 +90,7 @@ class TestMergeFunctions:
|
|||||||
def test_merge_mounts_mode_conflict(self) -> None:
|
def test_merge_mounts_mode_conflict(self) -> None:
|
||||||
"""Test that mount mode conflicts are resolved (later wins)."""
|
"""Test that mount mode conflicts are resolved (later wins)."""
|
||||||
from src.services.config_profile_resolver import ResolvedMount
|
from src.services.config_profile_resolver import ResolvedMount
|
||||||
|
|
||||||
overrides = {}
|
overrides = {}
|
||||||
result = _merge_mounts(
|
result = _merge_mounts(
|
||||||
{"/app": ResolvedMount(target="/app", mode="rw", files={})},
|
{"/app": ResolvedMount(target="/app", mode="rw", files={})},
|
||||||
@@ -97,6 +101,127 @@ class TestMergeFunctions:
|
|||||||
assert result["/app"].mode == "ro"
|
assert result["/app"].mode == "ro"
|
||||||
assert overrides == {"/app": "source"}
|
assert overrides == {"/app": "source"}
|
||||||
|
|
||||||
|
def test_merge_git_mounts_basic(self) -> None:
|
||||||
|
"""Test basic git mount merging normalizes to mappings format."""
|
||||||
|
result = _merge_git_mounts(
|
||||||
|
[],
|
||||||
|
[
|
||||||
|
{
|
||||||
|
"remote_url": "https://github.com/user/repo1.git",
|
||||||
|
"source_path": ".",
|
||||||
|
"target_path": "/app",
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"source",
|
||||||
|
)
|
||||||
|
assert len(result) == 1
|
||||||
|
assert result[0]["remote_url"] == "https://github.com/user/repo1.git"
|
||||||
|
assert "mappings" in result[0]
|
||||||
|
assert result[0]["mappings"] == [{"source_path": ".", "target_path": "/app"}]
|
||||||
|
|
||||||
|
def test_merge_git_mounts_concatenate_same_repo_branch(self) -> None:
|
||||||
|
"""Test that git mounts with same repo+branch concatenate mappings."""
|
||||||
|
result = _merge_git_mounts(
|
||||||
|
[
|
||||||
|
{
|
||||||
|
"remote_url": "https://github.com/user/repo1.git",
|
||||||
|
"source_path": ".",
|
||||||
|
"target_path": "/app",
|
||||||
|
"branch": "main",
|
||||||
|
}
|
||||||
|
],
|
||||||
|
[
|
||||||
|
{
|
||||||
|
"remote_url": "https://github.com/user/repo1.git",
|
||||||
|
"source_path": "src",
|
||||||
|
"target_path": "/src",
|
||||||
|
"branch": "main",
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"source",
|
||||||
|
)
|
||||||
|
assert len(result) == 1
|
||||||
|
assert result[0]["branch"] == "main"
|
||||||
|
mappings: list[dict[str, str]] = result[0]["mappings"]
|
||||||
|
assert len(mappings) == 2
|
||||||
|
assert {"source_path": ".", "target_path": "/app"} in mappings
|
||||||
|
assert {"source_path": "src", "target_path": "/src"} in mappings
|
||||||
|
|
||||||
|
def test_merge_git_mounts_dedup_same_mapping(self) -> None:
|
||||||
|
"""Test that duplicate mappings are deduplicated."""
|
||||||
|
result = _merge_git_mounts(
|
||||||
|
[
|
||||||
|
{
|
||||||
|
"remote_url": "https://github.com/user/repo1.git",
|
||||||
|
"source_path": ".",
|
||||||
|
"target_path": "/app",
|
||||||
|
"branch": "main",
|
||||||
|
}
|
||||||
|
],
|
||||||
|
[
|
||||||
|
{
|
||||||
|
"remote_url": "https://github.com/user/repo1.git",
|
||||||
|
"source_path": ".",
|
||||||
|
"target_path": "/app",
|
||||||
|
"branch": "main",
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"source",
|
||||||
|
)
|
||||||
|
assert len(result) == 1
|
||||||
|
assert len(result[0]["mappings"]) == 1
|
||||||
|
|
||||||
|
def test_merge_git_mounts_different_repos(self) -> None:
|
||||||
|
"""Test that git mounts with different repos are preserved."""
|
||||||
|
result = _merge_git_mounts(
|
||||||
|
[
|
||||||
|
{
|
||||||
|
"remote_url": "https://github.com/user/repo1.git",
|
||||||
|
"source_path": ".",
|
||||||
|
"target_path": "/app",
|
||||||
|
}
|
||||||
|
],
|
||||||
|
[
|
||||||
|
{
|
||||||
|
"remote_url": "https://github.com/user/repo2.git",
|
||||||
|
"source_path": ".",
|
||||||
|
"target_path": "/config",
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"source",
|
||||||
|
)
|
||||||
|
assert len(result) == 2
|
||||||
|
urls = {m["remote_url"] for m in result}
|
||||||
|
assert urls == {
|
||||||
|
"https://github.com/user/repo1.git",
|
||||||
|
"https://github.com/user/repo2.git",
|
||||||
|
}
|
||||||
|
|
||||||
|
def test_merge_git_mounts_different_branches(self) -> None:
|
||||||
|
"""Test that same repo with different branches are kept separate."""
|
||||||
|
result = _merge_git_mounts(
|
||||||
|
[
|
||||||
|
{
|
||||||
|
"remote_url": "https://github.com/user/repo1.git",
|
||||||
|
"source_path": ".",
|
||||||
|
"target_path": "/app",
|
||||||
|
"branch": "main",
|
||||||
|
}
|
||||||
|
],
|
||||||
|
[
|
||||||
|
{
|
||||||
|
"remote_url": "https://github.com/user/repo1.git",
|
||||||
|
"source_path": ".",
|
||||||
|
"target_path": "/app",
|
||||||
|
"branch": "dev",
|
||||||
|
}
|
||||||
|
],
|
||||||
|
"source",
|
||||||
|
)
|
||||||
|
assert len(result) == 2
|
||||||
|
branches = {m.get("branch") for m in result}
|
||||||
|
assert branches == {"main", "dev"}
|
||||||
|
|
||||||
|
|
||||||
class TestResolveProfile:
|
class TestResolveProfile:
|
||||||
"""Unit tests for profile resolution."""
|
"""Unit tests for profile resolution."""
|
||||||
@@ -124,7 +249,9 @@ class TestResolveProfile:
|
|||||||
assert result.files == {"test.txt": "content"}
|
assert result.files == {"test.txt": "content"}
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_resolve_profile_with_includes(self, db_session: AsyncSession) -> None:
|
async def test_resolve_profile_with_includes(
|
||||||
|
self, db_session: AsyncSession
|
||||||
|
) -> None:
|
||||||
"""Test resolving a profile that includes another."""
|
"""Test resolving a profile that includes another."""
|
||||||
user_id = uuid.uuid4()
|
user_id = uuid.uuid4()
|
||||||
|
|
||||||
@@ -168,7 +295,9 @@ class TestResolveProfile:
|
|||||||
assert result.included_profiles[0]["name"] == "base"
|
assert result.included_profiles[0]["name"] == "base"
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_resolve_profile_child_overrides_parent(self, db_session: AsyncSession) -> None:
|
async def test_resolve_profile_child_overrides_parent(
|
||||||
|
self, db_session: AsyncSession
|
||||||
|
) -> None:
|
||||||
"""Test that child profile values override parent values."""
|
"""Test that child profile values override parent values."""
|
||||||
user_id = uuid.uuid4()
|
user_id = uuid.uuid4()
|
||||||
|
|
||||||
@@ -205,7 +334,9 @@ class TestResolveProfile:
|
|||||||
assert result.env_overrides == {"VAR": "child"}
|
assert result.env_overrides == {"VAR": "child"}
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_resolve_profile_cycle_detection(self, db_session: AsyncSession) -> None:
|
async def test_resolve_profile_cycle_detection(
|
||||||
|
self, db_session: AsyncSession
|
||||||
|
) -> None:
|
||||||
"""Test that cycles are detected during resolution."""
|
"""Test that cycles are detected during resolution."""
|
||||||
user_id = uuid.uuid4()
|
user_id = uuid.uuid4()
|
||||||
|
|
||||||
@@ -250,6 +381,100 @@ class TestResolveProfile:
|
|||||||
with pytest.raises(ConfigProfileCycleError):
|
with pytest.raises(ConfigProfileCycleError):
|
||||||
await resolve_profile(db_session, profile_a.id)
|
await resolve_profile(db_session, profile_a.id)
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_resolve_profile_with_git_mounts(
|
||||||
|
self, db_session: AsyncSession
|
||||||
|
) -> None:
|
||||||
|
"""Test resolving a profile with git mounts normalizes to mappings."""
|
||||||
|
user_id = uuid.uuid4()
|
||||||
|
|
||||||
|
profile = ConfigProfile(
|
||||||
|
id=uuid.uuid4(),
|
||||||
|
user_id=user_id,
|
||||||
|
name="with-git-mounts",
|
||||||
|
env_vars={},
|
||||||
|
files={},
|
||||||
|
git_mounts=[
|
||||||
|
{
|
||||||
|
"remote_url": "https://github.com/user/repo1.git",
|
||||||
|
"source_path": ".",
|
||||||
|
"target_path": "/app",
|
||||||
|
},
|
||||||
|
],
|
||||||
|
)
|
||||||
|
db_session.add(profile)
|
||||||
|
await db_session.commit()
|
||||||
|
|
||||||
|
result = await resolve_profile(db_session, profile.id)
|
||||||
|
assert len(result.git_mounts) == 1
|
||||||
|
assert result.git_mounts[0]["remote_url"] == "https://github.com/user/repo1.git"
|
||||||
|
assert "mappings" in result.git_mounts[0]
|
||||||
|
assert result.git_mounts[0]["mappings"] == [
|
||||||
|
{"source_path": ".", "target_path": "/app"}
|
||||||
|
]
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_resolve_profile_with_git_mount_includes(
|
||||||
|
self, db_session: AsyncSession
|
||||||
|
) -> None:
|
||||||
|
"""Test resolving a profile that includes another with git mounts."""
|
||||||
|
user_id = uuid.uuid4()
|
||||||
|
|
||||||
|
# Create base profile with git mount
|
||||||
|
base = ConfigProfile(
|
||||||
|
id=uuid.uuid4(),
|
||||||
|
user_id=user_id,
|
||||||
|
name="base",
|
||||||
|
env_vars={},
|
||||||
|
files={},
|
||||||
|
git_mounts=[
|
||||||
|
{
|
||||||
|
"remote_url": "https://github.com/user/repo1.git",
|
||||||
|
"source_path": ".",
|
||||||
|
"target_path": "/app",
|
||||||
|
},
|
||||||
|
],
|
||||||
|
)
|
||||||
|
db_session.add(base)
|
||||||
|
|
||||||
|
# Create child profile with its own git mount
|
||||||
|
child = ConfigProfile(
|
||||||
|
id=uuid.uuid4(),
|
||||||
|
user_id=user_id,
|
||||||
|
name="child",
|
||||||
|
env_vars={},
|
||||||
|
files={},
|
||||||
|
git_mounts=[
|
||||||
|
{
|
||||||
|
"remote_url": "https://github.com/user/repo2.git",
|
||||||
|
"source_path": "config",
|
||||||
|
"target_path": "/config",
|
||||||
|
},
|
||||||
|
],
|
||||||
|
)
|
||||||
|
db_session.add(child)
|
||||||
|
await db_session.commit()
|
||||||
|
|
||||||
|
# Create include relationship
|
||||||
|
include = ConfigProfileInclude(
|
||||||
|
id=uuid.uuid4(),
|
||||||
|
profile_id=child.id,
|
||||||
|
included_profile_id=base.id,
|
||||||
|
order_index=0,
|
||||||
|
)
|
||||||
|
db_session.add(include)
|
||||||
|
await db_session.commit()
|
||||||
|
|
||||||
|
result = await resolve_profile(db_session, child.id)
|
||||||
|
assert len(result.git_mounts) == 2
|
||||||
|
urls = {m["remote_url"] for m in result.git_mounts}
|
||||||
|
assert urls == {
|
||||||
|
"https://github.com/user/repo1.git",
|
||||||
|
"https://github.com/user/repo2.git",
|
||||||
|
}
|
||||||
|
for m in result.git_mounts:
|
||||||
|
assert "mappings" in m
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_resolve_profile_not_found(self, db_session: AsyncSession) -> None:
|
async def test_resolve_profile_not_found(self, db_session: AsyncSession) -> None:
|
||||||
"""Test resolving a non-existent profile."""
|
"""Test resolving a non-existent profile."""
|
||||||
@@ -257,6 +482,82 @@ class TestResolveProfile:
|
|||||||
await resolve_profile(db_session, uuid.uuid4())
|
await resolve_profile(db_session, uuid.uuid4())
|
||||||
|
|
||||||
|
|
||||||
|
class TestApplyResolvedProfile:
|
||||||
|
"""Unit tests for apply_resolved_profile file-level mount behavior."""
|
||||||
|
|
||||||
|
def test_mounts_individual_files_not_directory(self, tmp_path) -> None:
|
||||||
|
"""Each file in a ResolvedMount should be mounted individually, not the staging dir."""
|
||||||
|
resolved = ResolvedProfile(
|
||||||
|
profile_id=uuid.uuid4(),
|
||||||
|
profile_name="test",
|
||||||
|
mounts={
|
||||||
|
"/app": ResolvedMount(
|
||||||
|
target="/app",
|
||||||
|
mode="rw",
|
||||||
|
files={
|
||||||
|
"config.json": '{"key": "value"}',
|
||||||
|
"nested/file.txt": "hello",
|
||||||
|
},
|
||||||
|
)
|
||||||
|
},
|
||||||
|
)
|
||||||
|
env, files, volumes, hints = apply_resolved_profile(str(tmp_path), resolved)
|
||||||
|
|
||||||
|
assert len(volumes) == 2
|
||||||
|
targets = {v["target"] for v in volumes}
|
||||||
|
assert "/app/config.json" in targets
|
||||||
|
assert "/app/nested/file.txt" in targets
|
||||||
|
# No directory-level mount
|
||||||
|
assert "/app" not in targets
|
||||||
|
|
||||||
|
def test_file_mount_preserves_sibling_files(self, tmp_path) -> None:
|
||||||
|
"""File-level mounts should not hide sibling files from other mounts."""
|
||||||
|
resolved = ResolvedProfile(
|
||||||
|
profile_id=uuid.uuid4(),
|
||||||
|
profile_name="test",
|
||||||
|
mounts={
|
||||||
|
"/workspace/x/y": ResolvedMount(
|
||||||
|
target="/workspace/x/y",
|
||||||
|
mode="rw",
|
||||||
|
files={"z.json": "override"},
|
||||||
|
)
|
||||||
|
},
|
||||||
|
)
|
||||||
|
env, files, volumes, hints = apply_resolved_profile(str(tmp_path), resolved)
|
||||||
|
|
||||||
|
assert len(volumes) == 1
|
||||||
|
assert volumes[0]["target"] == "/workspace/x/y/z.json"
|
||||||
|
assert volumes[0]["source"].endswith("z.json")
|
||||||
|
|
||||||
|
def test_empty_mount_produces_no_volumes(self, tmp_path) -> None:
|
||||||
|
"""A mount with no files should not produce any volume entries."""
|
||||||
|
resolved = ResolvedProfile(
|
||||||
|
profile_id=uuid.uuid4(),
|
||||||
|
profile_name="test",
|
||||||
|
mounts={"/app": ResolvedMount(target="/app", mode="rw", files={})},
|
||||||
|
)
|
||||||
|
env, files, volumes, hints = apply_resolved_profile(str(tmp_path), resolved)
|
||||||
|
assert volumes == []
|
||||||
|
|
||||||
|
def test_home_expansion_in_file_mount_target(self, tmp_path) -> None:
|
||||||
|
"""~ in mount target should be expanded to home_dir for file mounts."""
|
||||||
|
resolved = ResolvedProfile(
|
||||||
|
profile_id=uuid.uuid4(),
|
||||||
|
profile_name="test",
|
||||||
|
mounts={
|
||||||
|
"~/.config": ResolvedMount(
|
||||||
|
target="~/.config",
|
||||||
|
mode="rw",
|
||||||
|
files={"app.toml": "setting = 1"},
|
||||||
|
)
|
||||||
|
},
|
||||||
|
)
|
||||||
|
env, files, volumes, hints = apply_resolved_profile(
|
||||||
|
str(tmp_path), resolved, home_dir="/home/user"
|
||||||
|
)
|
||||||
|
assert volumes[0]["target"] == "/home/user/.config/app.toml"
|
||||||
|
|
||||||
|
|
||||||
class TestCheckIncludeCycle:
|
class TestCheckIncludeCycle:
|
||||||
"""Unit tests for include cycle checking."""
|
"""Unit tests for include cycle checking."""
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,112 @@
|
|||||||
|
"""Unit tests for docker service utilities."""
|
||||||
|
|
||||||
|
from unittest.mock import MagicMock, patch
|
||||||
|
|
||||||
|
import logging
|
||||||
|
|
||||||
|
from src.services.docker import (
|
||||||
|
get_container_id,
|
||||||
|
get_container_name,
|
||||||
|
sort_volumes_by_specificity,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestGetContainerId:
|
||||||
|
"""Tests for get_container_id."""
|
||||||
|
|
||||||
|
@patch("subprocess.run")
|
||||||
|
def test_lowercases_name_for_filter(self, mock_run) -> None:
|
||||||
|
"""Docker ps name filter is case-sensitive; we must lowercase."""
|
||||||
|
mock_run.return_value = MagicMock(returncode=0, stdout="abc123\n")
|
||||||
|
|
||||||
|
result = get_container_id("MyContainer-ABC")
|
||||||
|
|
||||||
|
assert result == "abc123"
|
||||||
|
call_args = mock_run.call_args[0][0]
|
||||||
|
# The filter must use lowercase
|
||||||
|
assert "name=mycontainer-abc" in call_args
|
||||||
|
|
||||||
|
@patch("subprocess.run")
|
||||||
|
def test_returns_none_when_not_found(self, mock_run) -> None:
|
||||||
|
mock_run.return_value = MagicMock(returncode=0, stdout="")
|
||||||
|
|
||||||
|
result = get_container_id("missing")
|
||||||
|
|
||||||
|
assert result is None
|
||||||
|
|
||||||
|
|
||||||
|
class TestGetContainerName:
|
||||||
|
"""Tests for get_container_name."""
|
||||||
|
|
||||||
|
@patch("subprocess.run")
|
||||||
|
def test_lowercases_name_for_filter(self, mock_run) -> None:
|
||||||
|
"""Docker ps name filter is case-sensitive; we must lowercase."""
|
||||||
|
mock_run.return_value = MagicMock(returncode=0, stdout="mycontainer-abc\n")
|
||||||
|
|
||||||
|
result = get_container_name("MyContainer-ABC")
|
||||||
|
|
||||||
|
assert result == "mycontainer-abc"
|
||||||
|
call_args = mock_run.call_args[0][0]
|
||||||
|
assert "name=mycontainer-abc" in call_args
|
||||||
|
|
||||||
|
@patch("subprocess.run")
|
||||||
|
def test_returns_none_when_not_found(self, mock_run) -> None:
|
||||||
|
mock_run.return_value = MagicMock(returncode=0, stdout="")
|
||||||
|
|
||||||
|
result = get_container_name("missing")
|
||||||
|
|
||||||
|
assert result is None
|
||||||
|
|
||||||
|
|
||||||
|
class TestSortVolumesBySpecificity:
|
||||||
|
"""Tests for sort_volumes_by_specificity."""
|
||||||
|
|
||||||
|
def test_parent_before_child(self) -> None:
|
||||||
|
"""A repo mount to /workspace/x should come before a file mount to /workspace/x/y/config.json."""
|
||||||
|
volumes = [
|
||||||
|
"/repo/x/y/config.json:/workspace/x/y/config.json",
|
||||||
|
"/repo/x:/workspace/x",
|
||||||
|
]
|
||||||
|
result = sort_volumes_by_specificity(volumes)
|
||||||
|
assert result[0] == "/repo/x:/workspace/x"
|
||||||
|
assert result[1] == "/repo/x/y/config.json:/workspace/x/y/config.json"
|
||||||
|
|
||||||
|
def test_stable_sort_for_equal_depth(self) -> None:
|
||||||
|
"""Mounts at the same depth preserve input order."""
|
||||||
|
volumes = [
|
||||||
|
"/a:/workspace/a",
|
||||||
|
"/b:/workspace/b",
|
||||||
|
"/c:/workspace/c",
|
||||||
|
]
|
||||||
|
result = sort_volumes_by_specificity(volumes)
|
||||||
|
assert result == volumes
|
||||||
|
|
||||||
|
def test_with_type_suffix(self) -> None:
|
||||||
|
"""Volume strings with :bind or :ro suffixes are parsed correctly."""
|
||||||
|
volumes = [
|
||||||
|
"/repo/x/y/config.json:/workspace/x/y/config.json:bind",
|
||||||
|
"/repo/x:/workspace/x:bind",
|
||||||
|
]
|
||||||
|
result = sort_volumes_by_specificity(volumes)
|
||||||
|
assert result[0] == "/repo/x:/workspace/x:bind"
|
||||||
|
assert result[1] == "/repo/x/y/config.json:/workspace/x/y/config.json:bind"
|
||||||
|
|
||||||
|
def test_empty_list(self) -> None:
|
||||||
|
"""Empty list returns empty list."""
|
||||||
|
assert sort_volumes_by_specificity([]) == []
|
||||||
|
|
||||||
|
def test_single_volume(self) -> None:
|
||||||
|
"""Single volume returns unchanged."""
|
||||||
|
volumes = ["/repo:/workspace"]
|
||||||
|
assert sort_volumes_by_specificity(volumes) == volumes
|
||||||
|
|
||||||
|
def test_duplicate_target_warning(self, caplog) -> None:
|
||||||
|
"""Duplicate targets trigger a warning."""
|
||||||
|
with caplog.at_level(logging.WARNING, logger="src.services.docker"):
|
||||||
|
volumes = [
|
||||||
|
"/a:/workspace/x",
|
||||||
|
"/b:/workspace/x",
|
||||||
|
]
|
||||||
|
sort_volumes_by_specificity(volumes)
|
||||||
|
assert "Duplicate mount targets detected" in caplog.text
|
||||||
|
assert "/workspace/x" in caplog.text
|
||||||
@@ -0,0 +1,148 @@
|
|||||||
|
"""Unit tests for InstanceEventBus."""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import uuid
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from src.services.event_bus import InstanceEventBus, InstanceEventPayload
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def event_bus() -> InstanceEventBus:
|
||||||
|
"""Provide a fresh EventBus instance with reset singleton state."""
|
||||||
|
bus = InstanceEventBus()
|
||||||
|
bus._reset_for_testing()
|
||||||
|
return bus
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def sample_payload() -> InstanceEventPayload:
|
||||||
|
"""Provide a sample event payload."""
|
||||||
|
return {
|
||||||
|
"event": "instance.started",
|
||||||
|
"instance_id": str(uuid.uuid4()),
|
||||||
|
"status": "starting",
|
||||||
|
"message": "Container starting...",
|
||||||
|
"metadata": {},
|
||||||
|
"timestamp": "2026-05-28T12:00:00Z",
|
||||||
|
"correlation_id": str(uuid.uuid4()),
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.unit
|
||||||
|
async def test_publish_delivers_to_all_subscribers(
|
||||||
|
event_bus: InstanceEventBus,
|
||||||
|
sample_payload: InstanceEventPayload,
|
||||||
|
) -> None:
|
||||||
|
"""All subscribed callbacks should receive the published payload."""
|
||||||
|
received: list[Any] = []
|
||||||
|
|
||||||
|
def callback_1(payload: InstanceEventPayload) -> None:
|
||||||
|
received.append(("callback_1", payload))
|
||||||
|
|
||||||
|
def callback_2(payload: InstanceEventPayload) -> None:
|
||||||
|
received.append(("callback_2", payload))
|
||||||
|
|
||||||
|
def callback_3(payload: InstanceEventPayload) -> None:
|
||||||
|
received.append(("callback_3", payload))
|
||||||
|
|
||||||
|
event_bus.subscribe("instance.started", callback_1)
|
||||||
|
event_bus.subscribe("instance.started", callback_2)
|
||||||
|
event_bus.subscribe("instance.started", callback_3)
|
||||||
|
|
||||||
|
await event_bus.publish("instance.started", sample_payload)
|
||||||
|
|
||||||
|
assert len(received) == 3
|
||||||
|
assert received[0][0] == "callback_1"
|
||||||
|
assert received[1][0] == "callback_2"
|
||||||
|
assert received[2][0] == "callback_3"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.unit
|
||||||
|
async def test_subscriber_exception_isolation(
|
||||||
|
event_bus: InstanceEventBus,
|
||||||
|
sample_payload: InstanceEventPayload,
|
||||||
|
) -> None:
|
||||||
|
"""If one subscriber raises, others should still receive the event."""
|
||||||
|
received: list[str] = []
|
||||||
|
|
||||||
|
def bad_callback(_payload: InstanceEventPayload) -> None:
|
||||||
|
raise RuntimeError("boom")
|
||||||
|
|
||||||
|
def good_callback(_payload: InstanceEventPayload) -> None:
|
||||||
|
received.append("good_callback")
|
||||||
|
|
||||||
|
event_bus.subscribe("instance.started", bad_callback)
|
||||||
|
event_bus.subscribe("instance.started", good_callback)
|
||||||
|
|
||||||
|
# Should not raise
|
||||||
|
await event_bus.publish("instance.started", sample_payload)
|
||||||
|
|
||||||
|
assert received == ["good_callback"]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.unit
|
||||||
|
async def test_unsubscribe_removes_callback(
|
||||||
|
event_bus: InstanceEventBus,
|
||||||
|
sample_payload: InstanceEventPayload,
|
||||||
|
) -> None:
|
||||||
|
"""After unsubscribing, the callback should not be called."""
|
||||||
|
received: list[str] = []
|
||||||
|
|
||||||
|
def callback(_payload: InstanceEventPayload) -> None:
|
||||||
|
received.append("callback")
|
||||||
|
|
||||||
|
unsubscribe = event_bus.subscribe("instance.started", callback)
|
||||||
|
unsubscribe()
|
||||||
|
|
||||||
|
await event_bus.publish("instance.started", sample_payload)
|
||||||
|
|
||||||
|
assert received == []
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.unit
|
||||||
|
async def test_publish_to_empty_subscriber_list(
|
||||||
|
event_bus: InstanceEventBus,
|
||||||
|
sample_payload: InstanceEventPayload,
|
||||||
|
) -> None:
|
||||||
|
"""Publishing to an event type with no subscribers should not raise."""
|
||||||
|
await event_bus.publish("instance.started", sample_payload)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.unit
|
||||||
|
async def test_async_subscriber_supported(
|
||||||
|
event_bus: InstanceEventBus,
|
||||||
|
sample_payload: InstanceEventPayload,
|
||||||
|
) -> None:
|
||||||
|
"""Async callbacks should be awaited correctly."""
|
||||||
|
received: list[str] = []
|
||||||
|
|
||||||
|
async def async_callback(_payload: InstanceEventPayload) -> None:
|
||||||
|
await asyncio.sleep(0)
|
||||||
|
received.append("async_callback")
|
||||||
|
|
||||||
|
event_bus.subscribe("instance.started", async_callback)
|
||||||
|
await event_bus.publish("instance.started", sample_payload)
|
||||||
|
|
||||||
|
assert received == ["async_callback"]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.unit
|
||||||
|
async def test_unsubscribe_all_clears_subscribers(
|
||||||
|
event_bus: InstanceEventBus,
|
||||||
|
sample_payload: InstanceEventPayload,
|
||||||
|
) -> None:
|
||||||
|
"""unsubscribe_all should remove all callbacks for an event type."""
|
||||||
|
received: list[str] = []
|
||||||
|
|
||||||
|
def callback(_payload: InstanceEventPayload) -> None:
|
||||||
|
received.append("callback")
|
||||||
|
|
||||||
|
event_bus.subscribe("instance.started", callback)
|
||||||
|
event_bus.unsubscribe_all("instance.started")
|
||||||
|
|
||||||
|
await event_bus.publish("instance.started", sample_payload)
|
||||||
|
|
||||||
|
assert received == []
|
||||||
@@ -0,0 +1,124 @@
|
|||||||
|
"""Unit tests for git mount resolution in tool instances."""
|
||||||
|
|
||||||
|
import os
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from src.api.tool_instances import (
|
||||||
|
_checkout_branch,
|
||||||
|
_expand_glob_source,
|
||||||
|
_resolve_single_git_mount,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestExpandGlobSource:
|
||||||
|
"""Unit tests for glob pattern expansion."""
|
||||||
|
|
||||||
|
def test_no_glob_single_file(self, tmp_path: Path) -> None:
|
||||||
|
"""Test non-glob path returns single file."""
|
||||||
|
test_file = tmp_path / "test.txt"
|
||||||
|
test_file.write_text("content")
|
||||||
|
|
||||||
|
result = _expand_glob_source(str(test_file), str(tmp_path))
|
||||||
|
assert len(result) == 1
|
||||||
|
assert result[0] == str(test_file)
|
||||||
|
|
||||||
|
def test_no_glob_missing_file(self, tmp_path: Path) -> None:
|
||||||
|
"""Test non-glob missing file returns empty list."""
|
||||||
|
missing_file = tmp_path / "missing.txt"
|
||||||
|
|
||||||
|
result = _expand_glob_source(str(missing_file), str(tmp_path))
|
||||||
|
assert len(result) == 0
|
||||||
|
|
||||||
|
def test_glob_pattern(self, tmp_path: Path) -> None:
|
||||||
|
"""Test glob pattern matches files."""
|
||||||
|
(tmp_path / "file1.txt").write_text("content1")
|
||||||
|
(tmp_path / "file2.txt").write_text("content2")
|
||||||
|
(tmp_path / "other.py").write_text("code")
|
||||||
|
|
||||||
|
result = _expand_glob_source(str(tmp_path / "*.txt"), str(tmp_path))
|
||||||
|
assert len(result) == 2
|
||||||
|
assert all(f.endswith(".txt") for f in result)
|
||||||
|
|
||||||
|
def test_glob_recursive(self, tmp_path: Path) -> None:
|
||||||
|
"""Test recursive glob pattern."""
|
||||||
|
subdir = tmp_path / "subdir"
|
||||||
|
subdir.mkdir()
|
||||||
|
(subdir / "nested.txt").write_text("content")
|
||||||
|
|
||||||
|
result = _expand_glob_source(str(tmp_path / "**" / "*.txt"), str(tmp_path))
|
||||||
|
assert len(result) == 1
|
||||||
|
assert "nested.txt" in result[0]
|
||||||
|
|
||||||
|
def test_glob_limit_enforced(self, tmp_path: Path) -> None:
|
||||||
|
"""Test that glob matches are limited to prevent abuse."""
|
||||||
|
# Create more than 100 files
|
||||||
|
for i in range(105):
|
||||||
|
(tmp_path / f"file{i}.txt").write_text("content")
|
||||||
|
|
||||||
|
result = _expand_glob_source(str(tmp_path / "*.txt"), str(tmp_path))
|
||||||
|
assert len(result) == 100 # MAX_GLOB_MATCHES limit
|
||||||
|
|
||||||
|
def test_glob_escapes_repo(self, tmp_path: Path) -> None:
|
||||||
|
"""Test that glob results outside repo are filtered."""
|
||||||
|
other_dir = tmp_path.parent / "other"
|
||||||
|
other_dir.mkdir(exist_ok=True)
|
||||||
|
(other_dir / "outside.txt").write_text("content")
|
||||||
|
|
||||||
|
result = _expand_glob_source(str(tmp_path.parent / "*" / "*.txt"), str(tmp_path))
|
||||||
|
# Should only include files within tmp_path, not other_dir
|
||||||
|
assert all(r.startswith(str(tmp_path)) for r in result)
|
||||||
|
|
||||||
|
|
||||||
|
class TestCheckoutBranch:
|
||||||
|
"""Unit tests for branch checkout."""
|
||||||
|
|
||||||
|
def test_checkout_existing_branch(self, tmp_path: Path) -> None:
|
||||||
|
"""Test checking out an existing branch."""
|
||||||
|
# Initialize git repo
|
||||||
|
os.system(f"cd {tmp_path} && git init && git config user.email 'test@test.com' && git config user.name 'Test'")
|
||||||
|
(tmp_path / "file.txt").write_text("content")
|
||||||
|
os.system(f"cd {tmp_path} && git add . && git commit -m 'initial'")
|
||||||
|
os.system(f"cd {tmp_path} && git branch feature")
|
||||||
|
|
||||||
|
_checkout_branch(str(tmp_path), "feature")
|
||||||
|
|
||||||
|
# Verify we're on feature branch
|
||||||
|
result = os.popen(f"cd {tmp_path} && git branch --show-current").read().strip()
|
||||||
|
assert result == "feature"
|
||||||
|
|
||||||
|
def test_checkout_nonexistent_branch(self, tmp_path: Path) -> None:
|
||||||
|
"""Test checking out a non-existent branch returns False."""
|
||||||
|
os.system(f"cd {tmp_path} && git init && git config user.email 'test@test.com' && git config user.name 'Test'")
|
||||||
|
(tmp_path / "file.txt").write_text("content")
|
||||||
|
os.system(f"cd {tmp_path} && git add . && git commit -m 'initial'")
|
||||||
|
|
||||||
|
result = _checkout_branch(str(tmp_path), "nonexistent")
|
||||||
|
assert result is False
|
||||||
|
|
||||||
|
|
||||||
|
class TestResolveSingleGitMount:
|
||||||
|
"""Unit tests for resolving a single git mount."""
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_resolve_missing_remote_url(self, db_session) -> None:
|
||||||
|
"""Test that missing remote_url returns empty list."""
|
||||||
|
git_mount = {
|
||||||
|
"source_path": ".",
|
||||||
|
"target_path": "/app",
|
||||||
|
}
|
||||||
|
|
||||||
|
result = await _resolve_single_git_mount(db_session, git_mount)
|
||||||
|
assert result == []
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_resolve_missing_target_path(self, db_session) -> None:
|
||||||
|
"""Test that missing target path returns empty list."""
|
||||||
|
git_mount = {
|
||||||
|
"remote_url": "https://github.com/user/repo.git",
|
||||||
|
"source_path": ".",
|
||||||
|
}
|
||||||
|
|
||||||
|
result = await _resolve_single_git_mount(db_session, git_mount)
|
||||||
|
assert result == []
|
||||||
@@ -0,0 +1,219 @@
|
|||||||
|
"""Unit tests for git mount resolution with multi-mapping support."""
|
||||||
|
|
||||||
|
import os
|
||||||
|
import tempfile
|
||||||
|
from unittest.mock import MagicMock, patch
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from src.api.tool_instances import (
|
||||||
|
_clone_git_repo,
|
||||||
|
_expand_glob_source,
|
||||||
|
_normalize_git_mount,
|
||||||
|
_resolve_git_mount_mappings,
|
||||||
|
_resolve_single_git_mount,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestNormalizeGitMount:
|
||||||
|
"""Tests for _normalize_git_mount."""
|
||||||
|
|
||||||
|
def test_legacy_to_mappings(self) -> None:
|
||||||
|
"""Legacy source_path + target_path becomes mappings array."""
|
||||||
|
entry = {
|
||||||
|
"remote_url": "https://github.com/user/repo.git",
|
||||||
|
"source_path": "packages/api",
|
||||||
|
"target_path": "/app/api",
|
||||||
|
"branch": "main",
|
||||||
|
}
|
||||||
|
result = _normalize_git_mount(entry)
|
||||||
|
assert "mappings" in result
|
||||||
|
assert result["mappings"] == [
|
||||||
|
{"source_path": "packages/api", "target_path": "/app/api"}
|
||||||
|
]
|
||||||
|
assert "source_path" not in result
|
||||||
|
assert "target_path" not in result
|
||||||
|
assert result["remote_url"] == "https://github.com/user/repo.git"
|
||||||
|
assert result["branch"] == "main"
|
||||||
|
|
||||||
|
def test_already_mappings(self) -> None:
|
||||||
|
"""Entry already with mappings is left unchanged."""
|
||||||
|
entry = {
|
||||||
|
"remote_url": "https://github.com/user/repo.git",
|
||||||
|
"branch": "main",
|
||||||
|
"mappings": [
|
||||||
|
{"source_path": "a", "target_path": "/a"},
|
||||||
|
{"source_path": "b", "target_path": "/b"},
|
||||||
|
],
|
||||||
|
}
|
||||||
|
result = _normalize_git_mount(entry)
|
||||||
|
assert result["mappings"] == [
|
||||||
|
{"source_path": "a", "target_path": "/a"},
|
||||||
|
{"source_path": "b", "target_path": "/b"},
|
||||||
|
]
|
||||||
|
assert "source_path" not in result
|
||||||
|
assert "target_path" not in result
|
||||||
|
|
||||||
|
def test_missing_target_path_no_mappings(self) -> None:
|
||||||
|
"""Entry with source_path but no target_path creates empty mappings."""
|
||||||
|
entry = {
|
||||||
|
"remote_url": "https://github.com/user/repo.git",
|
||||||
|
"source_path": "src",
|
||||||
|
}
|
||||||
|
result = _normalize_git_mount(entry)
|
||||||
|
assert "mappings" not in result
|
||||||
|
|
||||||
|
|
||||||
|
class TestResolveGitMountMappings:
|
||||||
|
"""Tests for _resolve_git_mount_mappings."""
|
||||||
|
|
||||||
|
def test_single_mapping(self) -> None:
|
||||||
|
"""A single mapping produces one volume mount."""
|
||||||
|
with tempfile.TemporaryDirectory() as repo_path:
|
||||||
|
os.makedirs(os.path.join(repo_path, "packages", "api"))
|
||||||
|
mappings = [
|
||||||
|
{"source_path": "packages/api", "target_path": "/app/api"},
|
||||||
|
]
|
||||||
|
result = _resolve_git_mount_mappings(repo_path, mappings, None)
|
||||||
|
assert len(result) == 1
|
||||||
|
assert result[0]["source"] == os.path.join(repo_path, "packages", "api")
|
||||||
|
assert result[0]["target"] == "/app/api"
|
||||||
|
assert result[0]["type"] == "bind"
|
||||||
|
|
||||||
|
def test_multiple_mappings(self) -> None:
|
||||||
|
"""Multiple mappings from same repo produce multiple mounts."""
|
||||||
|
with tempfile.TemporaryDirectory() as repo_path:
|
||||||
|
os.makedirs(os.path.join(repo_path, "packages", "api"))
|
||||||
|
os.makedirs(os.path.join(repo_path, "packages", "web"))
|
||||||
|
mappings = [
|
||||||
|
{"source_path": "packages/api", "target_path": "/app/api"},
|
||||||
|
{"source_path": "packages/web", "target_path": "/app/web"},
|
||||||
|
]
|
||||||
|
result = _resolve_git_mount_mappings(repo_path, mappings, None)
|
||||||
|
assert len(result) == 2
|
||||||
|
targets = {r["target"] for r in result}
|
||||||
|
assert targets == {"/app/api", "/app/web"}
|
||||||
|
|
||||||
|
def test_relative_target_path(self) -> None:
|
||||||
|
"""Relative target_path is resolved against working_directory."""
|
||||||
|
with tempfile.TemporaryDirectory() as repo_path:
|
||||||
|
os.makedirs(os.path.join(repo_path, "src"))
|
||||||
|
mappings = [
|
||||||
|
{"source_path": "src", "target_path": "code"},
|
||||||
|
]
|
||||||
|
result = _resolve_git_mount_mappings(repo_path, mappings, "/workspace")
|
||||||
|
assert len(result) == 1
|
||||||
|
assert result[0]["target"] == "/workspace/code"
|
||||||
|
|
||||||
|
def test_glob_expansion(self) -> None:
|
||||||
|
"""Glob patterns in source_path are expanded."""
|
||||||
|
with tempfile.TemporaryDirectory() as repo_path:
|
||||||
|
os.makedirs(os.path.join(repo_path, "packages", "api"))
|
||||||
|
os.makedirs(os.path.join(repo_path, "packages", "web"))
|
||||||
|
mappings = [
|
||||||
|
{"source_path": "packages/*", "target_path": "/app/packages"},
|
||||||
|
]
|
||||||
|
result = _resolve_git_mount_mappings(repo_path, mappings, None)
|
||||||
|
assert len(result) == 2
|
||||||
|
targets = {r["target"] for r in result}
|
||||||
|
assert targets == {
|
||||||
|
os.path.join("/app/packages", "packages", "api"),
|
||||||
|
os.path.join("/app/packages", "packages", "web"),
|
||||||
|
}
|
||||||
|
|
||||||
|
def test_missing_target_path_skipped(self) -> None:
|
||||||
|
"""Mapping without target_path is skipped."""
|
||||||
|
with tempfile.TemporaryDirectory() as repo_path:
|
||||||
|
mappings = [
|
||||||
|
{"source_path": "src"},
|
||||||
|
]
|
||||||
|
result = _resolve_git_mount_mappings(repo_path, mappings, None)
|
||||||
|
assert len(result) == 0
|
||||||
|
|
||||||
|
def test_no_working_directory_for_relative_target(self) -> None:
|
||||||
|
"""Relative target without working_directory is skipped."""
|
||||||
|
with tempfile.TemporaryDirectory() as repo_path:
|
||||||
|
os.makedirs(os.path.join(repo_path, "src"))
|
||||||
|
mappings = [
|
||||||
|
{"source_path": "src", "target_path": "code"},
|
||||||
|
]
|
||||||
|
result = _resolve_git_mount_mappings(repo_path, mappings, None)
|
||||||
|
assert len(result) == 0
|
||||||
|
|
||||||
|
|
||||||
|
class TestResolveSingleGitMount:
|
||||||
|
"""Tests for _resolve_single_git_mount."""
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_missing_remote_url(self) -> None:
|
||||||
|
"""Git mount without remote_url returns empty list."""
|
||||||
|
result = await _resolve_single_git_mount(
|
||||||
|
MagicMock(),
|
||||||
|
{"mappings": [{"source_path": ".", "target_path": "/app"}]},
|
||||||
|
"/tmp",
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
assert result == []
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_missing_instance_dir(self) -> None:
|
||||||
|
"""Git mount without instance_dir returns empty list."""
|
||||||
|
result = await _resolve_single_git_mount(
|
||||||
|
MagicMock(),
|
||||||
|
{
|
||||||
|
"remote_url": "https://github.com/user/repo.git",
|
||||||
|
"mappings": [{"source_path": ".", "target_path": "/app"}],
|
||||||
|
},
|
||||||
|
None,
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
assert result == []
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_legacy_format_normalized(self) -> None:
|
||||||
|
"""Legacy format is normalized and resolved."""
|
||||||
|
with tempfile.TemporaryDirectory() as instance_dir:
|
||||||
|
with patch(
|
||||||
|
"src.api.tool_instances._clone_git_repo",
|
||||||
|
return_value=os.path.join(instance_dir, "repo-clone"),
|
||||||
|
):
|
||||||
|
os.makedirs(os.path.join(instance_dir, "repo-clone", "src"))
|
||||||
|
result = await _resolve_single_git_mount(
|
||||||
|
MagicMock(),
|
||||||
|
{
|
||||||
|
"remote_url": "https://github.com/user/repo.git",
|
||||||
|
"source_path": "src",
|
||||||
|
"target_path": "/app/src",
|
||||||
|
},
|
||||||
|
instance_dir,
|
||||||
|
None,
|
||||||
|
)
|
||||||
|
assert len(result) == 1
|
||||||
|
assert result[0]["target"] == "/app/src"
|
||||||
|
|
||||||
|
|
||||||
|
class TestExpandGlobSource:
|
||||||
|
"""Tests for _expand_glob_source."""
|
||||||
|
|
||||||
|
def test_no_glob(self) -> None:
|
||||||
|
"""Non-glob path returns single item if exists."""
|
||||||
|
with tempfile.TemporaryDirectory() as tmp:
|
||||||
|
path = os.path.join(tmp, "file.txt")
|
||||||
|
open(path, "w").close()
|
||||||
|
result = _expand_glob_source(path, tmp)
|
||||||
|
assert result == [path]
|
||||||
|
|
||||||
|
def test_no_glob_missing(self) -> None:
|
||||||
|
"""Non-glob path that doesn't exist returns empty list."""
|
||||||
|
with tempfile.TemporaryDirectory() as tmp:
|
||||||
|
path = os.path.join(tmp, "missing.txt")
|
||||||
|
result = _expand_glob_source(path, tmp)
|
||||||
|
assert result == []
|
||||||
|
|
||||||
|
def test_glob_pattern(self) -> None:
|
||||||
|
"""Glob pattern expands to matched paths."""
|
||||||
|
with tempfile.TemporaryDirectory() as tmp:
|
||||||
|
open(os.path.join(tmp, "a.txt"), "w").close()
|
||||||
|
open(os.path.join(tmp, "b.txt"), "w").close()
|
||||||
|
result = _expand_glob_source(os.path.join(tmp, "*.txt"), tmp)
|
||||||
|
assert len(result) == 2
|
||||||
@@ -1,6 +1,5 @@
|
|||||||
"""Tests for git URL parsing utilities."""
|
"""Tests for git URL parsing utilities."""
|
||||||
|
|
||||||
import pytest
|
|
||||||
|
|
||||||
from src.utils.git_url_parser import extract_base_repo_url, is_valid_clone_url, parse_git_url
|
from src.utils.git_url_parser import extract_base_repo_url, is_valid_clone_url, parse_git_url
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,292 @@
|
|||||||
|
"""Unit tests for HealthMonitor state-transition logic."""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import uuid
|
||||||
|
from contextlib import suppress
|
||||||
|
from unittest.mock import patch
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
from sqlalchemy import select
|
||||||
|
|
||||||
|
from src.models.health_check import HealthCheck
|
||||||
|
from src.models.tool_instance import ToolInstance
|
||||||
|
from src.models.user import User
|
||||||
|
from src.services.event_bus import InstanceEventBus, InstanceEventPayload
|
||||||
|
from src.services.health_monitor import HealthMonitor, HealthSnapshot
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def event_bus() -> InstanceEventBus:
|
||||||
|
"""Provide a fresh EventBus instance."""
|
||||||
|
bus = InstanceEventBus()
|
||||||
|
bus._reset_for_testing()
|
||||||
|
return bus
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def health_monitor(event_bus: InstanceEventBus) -> HealthMonitor:
|
||||||
|
"""Provide a HealthMonitor with a short poll interval for testing."""
|
||||||
|
monitor = HealthMonitor(event_bus)
|
||||||
|
monitor.POLL_INTERVAL_SECONDS = 0.1
|
||||||
|
return monitor
|
||||||
|
|
||||||
|
|
||||||
|
async def _create_running_instance(db_session) -> ToolInstance:
|
||||||
|
"""Helper to create a user and a running tool instance."""
|
||||||
|
user = User(
|
||||||
|
id=uuid.uuid4(),
|
||||||
|
email="hm@example.com",
|
||||||
|
name="HM Test",
|
||||||
|
authentik_id="auth-hm",
|
||||||
|
)
|
||||||
|
db_session.add(user)
|
||||||
|
await db_session.commit()
|
||||||
|
|
||||||
|
instance = ToolInstance(
|
||||||
|
id=uuid.uuid4(),
|
||||||
|
name="hm-test-instance",
|
||||||
|
display_name="HM Test Instance",
|
||||||
|
tool_type_id=uuid.uuid4(),
|
||||||
|
repository_id=uuid.uuid4(),
|
||||||
|
project_id=uuid.uuid4(),
|
||||||
|
owner_id=user.id,
|
||||||
|
status="running",
|
||||||
|
container_id="container123",
|
||||||
|
public_url="https://example.trycloudflare.com",
|
||||||
|
)
|
||||||
|
db_session.add(instance)
|
||||||
|
await db_session.commit()
|
||||||
|
return instance
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.unit
|
||||||
|
async def test_detects_container_crash(
|
||||||
|
db_session,
|
||||||
|
event_bus: InstanceEventBus,
|
||||||
|
health_monitor: HealthMonitor,
|
||||||
|
) -> None:
|
||||||
|
"""Monitor should detect exited container and publish error event."""
|
||||||
|
instance = await _create_running_instance(db_session)
|
||||||
|
|
||||||
|
events_captured: list[InstanceEventPayload] = []
|
||||||
|
|
||||||
|
def capture_event(payload: InstanceEventPayload) -> None:
|
||||||
|
events_captured.append(payload)
|
||||||
|
|
||||||
|
event_bus.subscribe("instance.error", capture_event)
|
||||||
|
|
||||||
|
with (
|
||||||
|
patch(
|
||||||
|
"src.services.health_monitor.get_container_status",
|
||||||
|
return_value={"status": "exited", "exit_code": 137, "health": None},
|
||||||
|
),
|
||||||
|
patch(
|
||||||
|
"src.services.health_monitor.check_tunnel_health",
|
||||||
|
return_value={"healthy": False, "tunnel_status": "not_applicable"},
|
||||||
|
),
|
||||||
|
):
|
||||||
|
await health_monitor._check_instance(db_session, instance)
|
||||||
|
|
||||||
|
# Refresh instance from DB
|
||||||
|
await db_session.refresh(instance)
|
||||||
|
assert instance.status == "error"
|
||||||
|
|
||||||
|
# Event published
|
||||||
|
assert len(events_captured) == 1
|
||||||
|
assert events_captured[0]["event"] == "instance.error"
|
||||||
|
assert events_captured[0]["status"] == "error"
|
||||||
|
assert events_captured[0]["metadata"]["exit_code"] == 137
|
||||||
|
|
||||||
|
# Health check row inserted
|
||||||
|
result = await db_session.execute(
|
||||||
|
select(HealthCheck).where(HealthCheck.instance_id == instance.id)
|
||||||
|
)
|
||||||
|
check = result.scalar_one()
|
||||||
|
assert check.container_status == "exited"
|
||||||
|
assert check.exit_code == 137
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.unit
|
||||||
|
async def test_detects_tunnel_failure(
|
||||||
|
db_session,
|
||||||
|
event_bus: InstanceEventBus,
|
||||||
|
health_monitor: HealthMonitor,
|
||||||
|
) -> None:
|
||||||
|
"""Monitor should detect tunnel failure and mark unhealthy."""
|
||||||
|
instance = await _create_running_instance(db_session)
|
||||||
|
|
||||||
|
events_captured: list[InstanceEventPayload] = []
|
||||||
|
|
||||||
|
def capture_event(payload: InstanceEventPayload) -> None:
|
||||||
|
events_captured.append(payload)
|
||||||
|
|
||||||
|
event_bus.subscribe("instance.health_changed", capture_event)
|
||||||
|
|
||||||
|
with (
|
||||||
|
patch(
|
||||||
|
"src.services.health_monitor.get_container_status",
|
||||||
|
return_value={"status": "running", "exit_code": None, "health": "healthy"},
|
||||||
|
),
|
||||||
|
patch(
|
||||||
|
"src.services.health_monitor.check_tunnel_health",
|
||||||
|
return_value={
|
||||||
|
"healthy": False,
|
||||||
|
"tunnel_status": "error_response",
|
||||||
|
"status_code": 502,
|
||||||
|
},
|
||||||
|
),
|
||||||
|
):
|
||||||
|
await health_monitor._check_instance(db_session, instance)
|
||||||
|
|
||||||
|
await db_session.refresh(instance)
|
||||||
|
assert instance.status == "unhealthy"
|
||||||
|
|
||||||
|
assert len(events_captured) == 1
|
||||||
|
assert events_captured[0]["event"] == "instance.health_changed"
|
||||||
|
assert events_captured[0]["status"] == "unhealthy"
|
||||||
|
assert events_captured[0]["metadata"]["previous_status"] == "running"
|
||||||
|
|
||||||
|
result = await db_session.execute(
|
||||||
|
select(HealthCheck).where(HealthCheck.instance_id == instance.id)
|
||||||
|
)
|
||||||
|
check = result.scalar_one()
|
||||||
|
assert check.tunnel_healthy is False
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.unit
|
||||||
|
async def test_detects_recovery(
|
||||||
|
db_session,
|
||||||
|
event_bus: InstanceEventBus,
|
||||||
|
health_monitor: HealthMonitor,
|
||||||
|
) -> None:
|
||||||
|
"""Monitor should detect recovery from unhealthy to running."""
|
||||||
|
instance = await _create_running_instance(db_session)
|
||||||
|
instance.status = "unhealthy"
|
||||||
|
await db_session.commit()
|
||||||
|
|
||||||
|
# Seed last known state as unhealthy
|
||||||
|
health_monitor._last_known_state[instance.id] = HealthSnapshot(
|
||||||
|
container_status="running",
|
||||||
|
container_healthy=None,
|
||||||
|
tunnel_healthy=False,
|
||||||
|
exit_code=None,
|
||||||
|
)
|
||||||
|
|
||||||
|
events_captured: list[InstanceEventPayload] = []
|
||||||
|
|
||||||
|
def capture_event(payload: InstanceEventPayload) -> None:
|
||||||
|
events_captured.append(payload)
|
||||||
|
|
||||||
|
event_bus.subscribe("instance.health_changed", capture_event)
|
||||||
|
|
||||||
|
with (
|
||||||
|
patch(
|
||||||
|
"src.services.health_monitor.get_container_status",
|
||||||
|
return_value={"status": "running", "exit_code": None, "health": None},
|
||||||
|
),
|
||||||
|
patch(
|
||||||
|
"src.services.health_monitor.check_tunnel_health",
|
||||||
|
return_value={
|
||||||
|
"healthy": True,
|
||||||
|
"tunnel_status": "healthy",
|
||||||
|
"status_code": 200,
|
||||||
|
},
|
||||||
|
),
|
||||||
|
):
|
||||||
|
await health_monitor._check_instance(db_session, instance)
|
||||||
|
|
||||||
|
await db_session.refresh(instance)
|
||||||
|
assert instance.status == "running"
|
||||||
|
|
||||||
|
assert len(events_captured) == 1
|
||||||
|
assert events_captured[0]["status"] == "running"
|
||||||
|
assert events_captured[0]["metadata"]["previous_status"] == "unhealthy"
|
||||||
|
|
||||||
|
result = await db_session.execute(
|
||||||
|
select(HealthCheck).where(HealthCheck.instance_id == instance.id)
|
||||||
|
)
|
||||||
|
check = result.scalar_one()
|
||||||
|
assert check.tunnel_healthy is True
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.unit
|
||||||
|
async def test_skips_writes_when_no_state_change(
|
||||||
|
db_session,
|
||||||
|
event_bus: InstanceEventBus,
|
||||||
|
health_monitor: HealthMonitor,
|
||||||
|
) -> None:
|
||||||
|
"""Two identical polls should result in only one health_checks row."""
|
||||||
|
instance = await _create_running_instance(db_session)
|
||||||
|
|
||||||
|
with (
|
||||||
|
patch(
|
||||||
|
"src.services.health_monitor.get_container_status",
|
||||||
|
return_value={"status": "running", "exit_code": None, "health": None},
|
||||||
|
),
|
||||||
|
patch(
|
||||||
|
"src.services.health_monitor.check_tunnel_health",
|
||||||
|
return_value={
|
||||||
|
"healthy": True,
|
||||||
|
"tunnel_status": "healthy",
|
||||||
|
"status_code": 200,
|
||||||
|
},
|
||||||
|
),
|
||||||
|
):
|
||||||
|
await health_monitor._check_instance(db_session, instance)
|
||||||
|
await health_monitor._check_instance(db_session, instance)
|
||||||
|
|
||||||
|
result = await db_session.execute(
|
||||||
|
select(HealthCheck).where(HealthCheck.instance_id == instance.id)
|
||||||
|
)
|
||||||
|
assert len(result.scalars().all()) == 1
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.unit
|
||||||
|
async def test_docker_exception_resilience(
|
||||||
|
db_session,
|
||||||
|
event_bus: InstanceEventBus,
|
||||||
|
health_monitor: HealthMonitor,
|
||||||
|
) -> None:
|
||||||
|
"""Docker exception should be caught and not propagate."""
|
||||||
|
instance = await _create_running_instance(db_session)
|
||||||
|
|
||||||
|
events_captured: list[InstanceEventPayload] = []
|
||||||
|
|
||||||
|
def capture_event(payload: InstanceEventPayload) -> None:
|
||||||
|
events_captured.append(payload)
|
||||||
|
|
||||||
|
event_bus.subscribe("instance.error", capture_event)
|
||||||
|
event_bus.subscribe("instance.health_changed", capture_event)
|
||||||
|
|
||||||
|
with patch(
|
||||||
|
"src.services.health_monitor.get_container_status",
|
||||||
|
side_effect=RuntimeError("docker exploded"),
|
||||||
|
):
|
||||||
|
# Should not raise
|
||||||
|
await health_monitor._check_instance(db_session, instance)
|
||||||
|
|
||||||
|
# No DB writes
|
||||||
|
result = await db_session.execute(
|
||||||
|
select(HealthCheck).where(HealthCheck.instance_id == instance.id)
|
||||||
|
)
|
||||||
|
assert result.scalar_one_or_none() is None
|
||||||
|
|
||||||
|
# No events published
|
||||||
|
assert events_captured == []
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.unit
|
||||||
|
async def test_monitor_start_stop(health_monitor: HealthMonitor) -> None:
|
||||||
|
"""Start and stop should manage the background task."""
|
||||||
|
health_monitor.start()
|
||||||
|
task = health_monitor._task
|
||||||
|
assert task is not None
|
||||||
|
assert not task.done()
|
||||||
|
|
||||||
|
health_monitor.stop()
|
||||||
|
if task is not None and not task.done():
|
||||||
|
with suppress(asyncio.CancelledError):
|
||||||
|
await task
|
||||||
|
assert task is not None
|
||||||
|
assert task.cancelled() or task.done()
|
||||||
|
assert health_monitor._last_known_state == {}
|
||||||
@@ -0,0 +1,110 @@
|
|||||||
|
"""Unit tests for ~ / $HOME expansion in container paths."""
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from src.api.tool_instances import _resolve_git_mount_mappings
|
||||||
|
from src.services.config_profile_resolver import expand_container_path
|
||||||
|
from src.services.manifest_compiler import get_manifest_home_dir
|
||||||
|
|
||||||
|
|
||||||
|
class TestExpandContainerPath:
|
||||||
|
"""Tests for expand_container_path helper."""
|
||||||
|
|
||||||
|
def test_tilde_slash_expands(self) -> None:
|
||||||
|
"""~/foo should expand to home_dir/foo."""
|
||||||
|
assert (
|
||||||
|
expand_container_path("~/workspace", "/home/user") == "/home/user/workspace"
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_tilde_alone_expands(self) -> None:
|
||||||
|
"""~ should expand to home_dir."""
|
||||||
|
assert expand_container_path("~", "/home/user") == "/home/user"
|
||||||
|
|
||||||
|
def test_dollar_home_slash_expands(self) -> None:
|
||||||
|
"""$HOME/foo should expand to home_dir/foo."""
|
||||||
|
assert (
|
||||||
|
expand_container_path("$HOME/workspace", "/home/user")
|
||||||
|
== "/home/user/workspace"
|
||||||
|
)
|
||||||
|
|
||||||
|
def test_dollar_home_alone_expands(self) -> None:
|
||||||
|
"""$HOME should expand to home_dir."""
|
||||||
|
assert expand_container_path("$HOME", "/home/user") == "/home/user"
|
||||||
|
|
||||||
|
def test_absolute_path_unchanged(self) -> None:
|
||||||
|
"""Absolute paths should not be modified."""
|
||||||
|
assert expand_container_path("/app/workspace", "/home/user") == "/app/workspace"
|
||||||
|
|
||||||
|
def test_relative_path_unchanged(self) -> None:
|
||||||
|
"""Relative paths should not be modified."""
|
||||||
|
assert expand_container_path("workspace", "/home/user") == "workspace"
|
||||||
|
|
||||||
|
def test_tilde_in_middle_unchanged(self) -> None:
|
||||||
|
"""~ in the middle of a path should not expand."""
|
||||||
|
assert expand_container_path("/app/~user", "/home/user") == "/app/~user"
|
||||||
|
|
||||||
|
def test_dollar_home_in_middle_unchanged(self) -> None:
|
||||||
|
"""$HOME in the middle of a path should not expand."""
|
||||||
|
assert expand_container_path("/app/$HOMEuser", "/home/user") == "/app/$HOMEuser"
|
||||||
|
|
||||||
|
def test_root_home(self) -> None:
|
||||||
|
"""Expansion works with /root as home."""
|
||||||
|
assert expand_container_path("~/config", "/root") == "/root/config"
|
||||||
|
|
||||||
|
|
||||||
|
class TestGetManifestHomeDir:
|
||||||
|
"""Tests for get_manifest_home_dir helper."""
|
||||||
|
|
||||||
|
def test_with_user_block(self) -> None:
|
||||||
|
"""Manifest with user block returns /home/{name}."""
|
||||||
|
manifest = {"user": {"name": "developer", "uid": 1000, "gid": 1000}}
|
||||||
|
assert get_manifest_home_dir(manifest) == "/home/developer"
|
||||||
|
|
||||||
|
def test_without_user_block(self) -> None:
|
||||||
|
"""Manifest without user block returns /root."""
|
||||||
|
manifest = {"base_image": "ubuntu:24.04"}
|
||||||
|
assert get_manifest_home_dir(manifest) == "/root"
|
||||||
|
|
||||||
|
def test_with_empty_user_name(self) -> None:
|
||||||
|
"""Manifest with empty user name returns /root."""
|
||||||
|
manifest = {"user": {"name": "", "uid": 1000, "gid": 1000}}
|
||||||
|
assert get_manifest_home_dir(manifest) == "/root"
|
||||||
|
|
||||||
|
def test_with_none_user_name(self) -> None:
|
||||||
|
"""Manifest with None user name returns /root."""
|
||||||
|
manifest = {"user": {"name": None, "uid": 1000, "gid": 1000}}
|
||||||
|
assert get_manifest_home_dir(manifest) == "/root"
|
||||||
|
|
||||||
|
|
||||||
|
class TestResolveGitMountMappingsExpansion:
|
||||||
|
"""Tests that git mount mapping targets expand ~ and $HOME."""
|
||||||
|
|
||||||
|
def test_tilde_target_expansion(self, tmp_path) -> None:
|
||||||
|
"""Mapping with ~/repo target expands to home dir."""
|
||||||
|
(tmp_path / "src").mkdir()
|
||||||
|
mappings = [{"source_path": "src", "target_path": "~/repo"}]
|
||||||
|
result = _resolve_git_mount_mappings(
|
||||||
|
str(tmp_path), mappings, None, "/home/user"
|
||||||
|
)
|
||||||
|
assert len(result) == 1
|
||||||
|
assert result[0]["target"] == "/home/user/repo"
|
||||||
|
|
||||||
|
def test_dollar_home_target_expansion(self, tmp_path) -> None:
|
||||||
|
"""Mapping with $HOME/repo target expands to home dir."""
|
||||||
|
(tmp_path / "src").mkdir()
|
||||||
|
mappings = [{"source_path": "src", "target_path": "$HOME/repo"}]
|
||||||
|
result = _resolve_git_mount_mappings(
|
||||||
|
str(tmp_path), mappings, None, "/home/user"
|
||||||
|
)
|
||||||
|
assert len(result) == 1
|
||||||
|
assert result[0]["target"] == "/home/user/repo"
|
||||||
|
|
||||||
|
def test_absolute_target_unchanged(self, tmp_path) -> None:
|
||||||
|
"""Absolute target paths are not modified."""
|
||||||
|
(tmp_path / "src").mkdir()
|
||||||
|
mappings = [{"source_path": "src", "target_path": "/app/src"}]
|
||||||
|
result = _resolve_git_mount_mappings(
|
||||||
|
str(tmp_path), mappings, None, "/home/user"
|
||||||
|
)
|
||||||
|
assert len(result) == 1
|
||||||
|
assert result[0]["target"] == "/app/src"
|
||||||
@@ -0,0 +1,49 @@
|
|||||||
|
"""Unit tests for lifecycle hook helpers."""
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from src.services.lifecycle_hooks import _derive_title, _should_notify
|
||||||
|
|
||||||
|
|
||||||
|
class TestDeriveTitle:
|
||||||
|
"""Tests for _derive_title."""
|
||||||
|
|
||||||
|
def test_known_event_types(self) -> None:
|
||||||
|
assert _derive_title("instance.created") == "Container created"
|
||||||
|
assert _derive_title("instance.started") == "Container started"
|
||||||
|
assert _derive_title("instance.stopped") == "Container stopped"
|
||||||
|
assert _derive_title("instance.restarted") == "Container restarted"
|
||||||
|
assert _derive_title("instance.deleted") == "Container deleted"
|
||||||
|
assert _derive_title("instance.error") == "Container error"
|
||||||
|
assert _derive_title("instance.health_changed") == "Container ready"
|
||||||
|
|
||||||
|
def test_unknown_event_type(self) -> None:
|
||||||
|
assert _derive_title("instance.custom_event") == "Custom Event"
|
||||||
|
|
||||||
|
|
||||||
|
class TestShouldNotify:
|
||||||
|
"""Tests for _should_notify filtering."""
|
||||||
|
|
||||||
|
def test_error_events_are_notified(self) -> None:
|
||||||
|
assert _should_notify("instance.error", "error") is True
|
||||||
|
assert _should_notify("instance.error", None) is True
|
||||||
|
|
||||||
|
def test_health_changed_running_is_notified(self) -> None:
|
||||||
|
assert _should_notify("instance.health_changed", "running") is True
|
||||||
|
|
||||||
|
def test_created_started_stopped_restarted_deleted_filtered(self) -> None:
|
||||||
|
for event in [
|
||||||
|
"instance.created",
|
||||||
|
"instance.started",
|
||||||
|
"instance.stopped",
|
||||||
|
"instance.restarted",
|
||||||
|
"instance.deleted",
|
||||||
|
]:
|
||||||
|
assert _should_notify(event, "pending") is False
|
||||||
|
assert _should_notify(event, "running") is False
|
||||||
|
assert _should_notify(event, None) is False
|
||||||
|
|
||||||
|
def test_health_changed_non_running_filtered(self) -> None:
|
||||||
|
assert _should_notify("instance.health_changed", "unhealthy") is False
|
||||||
|
assert _should_notify("instance.health_changed", "starting") is False
|
||||||
|
assert _should_notify("instance.health_changed", None) is False
|
||||||
@@ -0,0 +1,365 @@
|
|||||||
|
"""Unit tests for the manifest compiler."""
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from src.services.manifest_compiler import (
|
||||||
|
compile_compose,
|
||||||
|
compile_dockerfile,
|
||||||
|
compile_entrypoint,
|
||||||
|
compute_image_tag,
|
||||||
|
deep_merge,
|
||||||
|
get_manifest_home_dir,
|
||||||
|
merge_with_config,
|
||||||
|
resolve_base,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class TestResolveBase:
|
||||||
|
"""Tests for resolve_base."""
|
||||||
|
|
||||||
|
def test_returns_manifest_unchanged_when_no_base(self) -> None:
|
||||||
|
manifest = {"name": "test", "base_image": "ubuntu:24.04"}
|
||||||
|
result = resolve_base(manifest)
|
||||||
|
assert result["name"] == "test"
|
||||||
|
assert "base_definition_id" not in result
|
||||||
|
|
||||||
|
|
||||||
|
class TestDeepMerge:
|
||||||
|
"""Tests for deep_merge."""
|
||||||
|
|
||||||
|
def test_packages_are_unioned(self) -> None:
|
||||||
|
base = {"packages": {"apt": ["curl", "git"]}}
|
||||||
|
override = {"packages": {"apt": ["neovim"]}}
|
||||||
|
result = deep_merge(base, override)
|
||||||
|
assert result["packages"]["apt"] == ["curl", "git", "neovim"]
|
||||||
|
|
||||||
|
def test_node_version_overrides(self) -> None:
|
||||||
|
base = {"packages": {"node": {"version": "18"}}}
|
||||||
|
override = {"packages": {"node": {"version": "20"}}}
|
||||||
|
result = deep_merge(base, override)
|
||||||
|
assert result["packages"]["node"]["version"] == "20"
|
||||||
|
|
||||||
|
def test_env_is_merged_with_override_winning(self) -> None:
|
||||||
|
base = {"env": {"FOO": "base", "BAR": "base"}}
|
||||||
|
override = {"env": {"FOO": "override"}}
|
||||||
|
result = deep_merge(base, override)
|
||||||
|
assert result["env"]["FOO"] == "override"
|
||||||
|
assert result["env"]["BAR"] == "base"
|
||||||
|
|
||||||
|
def test_build_scripts_are_concatenated(self) -> None:
|
||||||
|
base = {"scripts": {"build": ["echo base"]}}
|
||||||
|
override = {"scripts": {"build": ["echo override"]}}
|
||||||
|
result = deep_merge(base, override)
|
||||||
|
assert result["scripts"]["build"] == ["echo base", "echo override"]
|
||||||
|
|
||||||
|
def test_mounts_are_concatenated(self) -> None:
|
||||||
|
base = {"mounts": [{"name": "base-mount", "target": "/base"}]}
|
||||||
|
override = {"mounts": [{"name": "tool-mount", "target": "/tool"}]}
|
||||||
|
result = deep_merge(base, override)
|
||||||
|
assert len(result["mounts"]) == 2
|
||||||
|
|
||||||
|
def test_user_is_overridden_entirely(self) -> None:
|
||||||
|
base = {"user": {"name": "base", "uid": 1000}}
|
||||||
|
override = {"user": {"name": "tool", "uid": 1001}}
|
||||||
|
result = deep_merge(base, override)
|
||||||
|
assert result["user"]["name"] == "tool"
|
||||||
|
assert result["user"]["uid"] == 1001
|
||||||
|
|
||||||
|
|
||||||
|
class TestCompileDockerfile:
|
||||||
|
"""Tests for compile_dockerfile."""
|
||||||
|
|
||||||
|
def test_includes_from(self) -> None:
|
||||||
|
manifest = {"base_image": "ubuntu:24.04", "name": "test"}
|
||||||
|
df = compile_dockerfile(manifest)
|
||||||
|
assert "FROM ubuntu:24.04" in df
|
||||||
|
|
||||||
|
def test_installs_apt_packages(self) -> None:
|
||||||
|
manifest = {
|
||||||
|
"base_image": "ubuntu:24.04",
|
||||||
|
"name": "test",
|
||||||
|
"packages": {"apt": ["curl", "git"]},
|
||||||
|
}
|
||||||
|
df = compile_dockerfile(manifest)
|
||||||
|
assert "apt-get install -y" in df
|
||||||
|
assert "curl" in df
|
||||||
|
assert "git" in df
|
||||||
|
assert "rm -rf /var/lib/apt/lists/*" in df
|
||||||
|
|
||||||
|
def test_installs_node(self) -> None:
|
||||||
|
manifest = {
|
||||||
|
"base_image": "ubuntu:24.04",
|
||||||
|
"name": "test",
|
||||||
|
"packages": {"node": {"version": "20"}},
|
||||||
|
}
|
||||||
|
df = compile_dockerfile(manifest)
|
||||||
|
assert "nodesource.com/setup_20.x" in df
|
||||||
|
|
||||||
|
def test_installs_npm_global(self) -> None:
|
||||||
|
manifest = {
|
||||||
|
"base_image": "ubuntu:24.04",
|
||||||
|
"name": "test",
|
||||||
|
"packages": {"npm_global": ["@scope/pkg"]},
|
||||||
|
}
|
||||||
|
df = compile_dockerfile(manifest)
|
||||||
|
assert "npm install -g @scope/pkg" in df
|
||||||
|
|
||||||
|
def test_creates_user(self) -> None:
|
||||||
|
manifest = {
|
||||||
|
"base_image": "ubuntu:24.04",
|
||||||
|
"name": "test",
|
||||||
|
"user": {"name": "dev", "uid": 1001, "gid": 1001},
|
||||||
|
}
|
||||||
|
df = compile_dockerfile(manifest)
|
||||||
|
assert "groupadd -g 1001 dev" in df
|
||||||
|
assert "useradd -u 1001 -g 1001" in df
|
||||||
|
assert "USER dev" in df
|
||||||
|
|
||||||
|
def test_build_scripts_as_run_commands(self) -> None:
|
||||||
|
manifest = {
|
||||||
|
"base_image": "ubuntu:24.04",
|
||||||
|
"name": "test",
|
||||||
|
"scripts": {"build": ["echo hello", "echo world"]},
|
||||||
|
}
|
||||||
|
df = compile_dockerfile(manifest)
|
||||||
|
assert "RUN echo hello" in df
|
||||||
|
assert "RUN echo world" in df
|
||||||
|
|
||||||
|
def test_creates_mount_directories(self) -> None:
|
||||||
|
manifest = {
|
||||||
|
"base_image": "ubuntu:24.04",
|
||||||
|
"name": "test",
|
||||||
|
"user": {"name": "dev", "uid": 1001, "gid": 1001},
|
||||||
|
"mounts": [
|
||||||
|
{"name": "ws", "target": "/workspace"},
|
||||||
|
{"name": "cfg", "target": "/config"},
|
||||||
|
],
|
||||||
|
}
|
||||||
|
df = compile_dockerfile(manifest)
|
||||||
|
assert "mkdir -p /workspace /config" in df
|
||||||
|
assert "chown -R dev:dev /workspace /config" in df
|
||||||
|
|
||||||
|
def test_entrypoint_for_startup_scripts(self) -> None:
|
||||||
|
manifest = {
|
||||||
|
"base_image": "ubuntu:24.04",
|
||||||
|
"name": "test",
|
||||||
|
"scripts": {"startup": ["echo start"]},
|
||||||
|
}
|
||||||
|
df = compile_dockerfile(manifest)
|
||||||
|
assert 'ENTRYPOINT ["/usr/local/bin/headquarter-entrypoint"]' in df
|
||||||
|
|
||||||
|
def test_cmd_from_runtime(self) -> None:
|
||||||
|
manifest = {
|
||||||
|
"base_image": "ubuntu:24.04",
|
||||||
|
"name": "test",
|
||||||
|
"runtime": {"command": ["/bin/bash", "-il"]},
|
||||||
|
}
|
||||||
|
df = compile_dockerfile(manifest)
|
||||||
|
assert 'CMD ["/bin/bash", "-il"]' in df
|
||||||
|
|
||||||
|
def test_default_cmd_when_no_runtime(self) -> None:
|
||||||
|
manifest = {"base_image": "ubuntu:24.04", "name": "test"}
|
||||||
|
df = compile_dockerfile(manifest)
|
||||||
|
assert 'CMD ["/bin/bash"]' in df
|
||||||
|
|
||||||
|
def test_sets_home_env_for_user(self) -> None:
|
||||||
|
manifest = {
|
||||||
|
"base_image": "ubuntu:24.04",
|
||||||
|
"name": "test",
|
||||||
|
"user": {"name": "dev", "uid": 1001, "gid": 1001},
|
||||||
|
}
|
||||||
|
df = compile_dockerfile(manifest)
|
||||||
|
assert "ENV HOME=/home/dev" in df
|
||||||
|
assert "ENV USER=dev" in df
|
||||||
|
|
||||||
|
def test_no_home_env_without_user(self) -> None:
|
||||||
|
manifest = {"base_image": "ubuntu:24.04", "name": "test"}
|
||||||
|
df = compile_dockerfile(manifest)
|
||||||
|
assert "ENV HOME=" not in df
|
||||||
|
assert "ENV USER=" not in df
|
||||||
|
|
||||||
|
|
||||||
|
class TestGetManifestHomeDir:
|
||||||
|
"""Tests for get_manifest_home_dir."""
|
||||||
|
|
||||||
|
def test_with_user_name(self) -> None:
|
||||||
|
manifest = {"user": {"name": "dev", "uid": 1001, "gid": 1001}}
|
||||||
|
assert get_manifest_home_dir(manifest) == "/home/dev"
|
||||||
|
|
||||||
|
def test_without_user(self) -> None:
|
||||||
|
manifest = {"base_image": "ubuntu:24.04"}
|
||||||
|
assert get_manifest_home_dir(manifest) == "/root"
|
||||||
|
|
||||||
|
def test_with_empty_user_name(self) -> None:
|
||||||
|
manifest = {"user": {"name": "", "uid": 1001, "gid": 1001}}
|
||||||
|
assert get_manifest_home_dir(manifest) == "/root"
|
||||||
|
|
||||||
|
|
||||||
|
class TestCompileEntrypoint:
|
||||||
|
"""Tests for compile_entrypoint."""
|
||||||
|
|
||||||
|
def test_includes_shebang_and_set_e(self) -> None:
|
||||||
|
manifest = {"scripts": {"startup": ["echo hello"]}}
|
||||||
|
ep = compile_entrypoint(manifest)
|
||||||
|
assert "#!/bin/bash" in ep
|
||||||
|
assert "set -e" in ep
|
||||||
|
|
||||||
|
def test_includes_startup_scripts(self) -> None:
|
||||||
|
manifest = {"scripts": {"startup": ["echo hello", "echo world"]}}
|
||||||
|
ep = compile_entrypoint(manifest)
|
||||||
|
assert "echo hello" in ep
|
||||||
|
assert "echo world" in ep
|
||||||
|
|
||||||
|
def test_ends_with_exec(self) -> None:
|
||||||
|
manifest: dict = {"scripts": {"startup": []}}
|
||||||
|
ep = compile_entrypoint(manifest)
|
||||||
|
assert 'exec "$@"' in ep
|
||||||
|
|
||||||
|
|
||||||
|
class TestCompileCompose:
|
||||||
|
"""Tests for compile_compose."""
|
||||||
|
|
||||||
|
def test_includes_image_and_container_name(self) -> None:
|
||||||
|
manifest = {"name": "test", "interface_type": "terminal"}
|
||||||
|
vars_dict = {"IMAGE_TAG": "test:v1", "INSTANCE_NAME": "test-1"}
|
||||||
|
compose = compile_compose(manifest, vars_dict)
|
||||||
|
assert "image: test:v1" in compose
|
||||||
|
assert "container_name: test-1" in compose
|
||||||
|
|
||||||
|
def test_terminal_fields(self) -> None:
|
||||||
|
manifest = {
|
||||||
|
"name": "test",
|
||||||
|
"interface_type": "terminal",
|
||||||
|
"runtime": {"stdin_open": True, "tty": True, "working_dir": "/workspace"},
|
||||||
|
}
|
||||||
|
compose = compile_compose(manifest, {"IMAGE_TAG": "t", "INSTANCE_NAME": "n"})
|
||||||
|
assert "stdin_open: true" in compose
|
||||||
|
assert "tty: true" in compose
|
||||||
|
assert "working_dir: /workspace" in compose
|
||||||
|
|
||||||
|
def test_web_ports(self) -> None:
|
||||||
|
manifest = {
|
||||||
|
"name": "test",
|
||||||
|
"interface_type": "web",
|
||||||
|
"default_port": 8080,
|
||||||
|
}
|
||||||
|
compose = compile_compose(
|
||||||
|
manifest, {"IMAGE_TAG": "t", "INSTANCE_NAME": "n", "TOOL_PORT": "3000"}
|
||||||
|
)
|
||||||
|
assert "3000:8080" in compose
|
||||||
|
|
||||||
|
def test_user_override(self) -> None:
|
||||||
|
manifest = {
|
||||||
|
"name": "test",
|
||||||
|
"interface_type": "terminal",
|
||||||
|
"user": {"uid": 1001, "gid": 1001},
|
||||||
|
}
|
||||||
|
compose = compile_compose(manifest, {"IMAGE_TAG": "t", "INSTANCE_NAME": "n"})
|
||||||
|
assert "user: 1001:1001" in compose
|
||||||
|
|
||||||
|
def test_mounts_resolved(self) -> None:
|
||||||
|
manifest = {
|
||||||
|
"name": "test",
|
||||||
|
"interface_type": "terminal",
|
||||||
|
"mounts": [
|
||||||
|
{"name": "ws", "target": "/workspace", "source_type": "repo"},
|
||||||
|
{
|
||||||
|
"name": "ssh",
|
||||||
|
"target": "/home/user/.ssh",
|
||||||
|
"source_type": "ssh_key",
|
||||||
|
"readonly": True,
|
||||||
|
},
|
||||||
|
],
|
||||||
|
}
|
||||||
|
compose = compile_compose(
|
||||||
|
manifest,
|
||||||
|
{
|
||||||
|
"IMAGE_TAG": "t",
|
||||||
|
"INSTANCE_NAME": "n",
|
||||||
|
"REPO_PATH": "/repos/myrepo",
|
||||||
|
"SSH_PATH": "/keys/ssh",
|
||||||
|
},
|
||||||
|
)
|
||||||
|
assert "/repos/myrepo:/workspace" in compose
|
||||||
|
assert "/keys/ssh:/home/user/.ssh:ro" in compose
|
||||||
|
|
||||||
|
def test_extra_volumes_appended(self) -> None:
|
||||||
|
manifest = {"name": "test", "interface_type": "terminal"}
|
||||||
|
compose = compile_compose(
|
||||||
|
manifest,
|
||||||
|
{
|
||||||
|
"IMAGE_TAG": "t",
|
||||||
|
"INSTANCE_NAME": "n",
|
||||||
|
"EXTRA_VOLUMES": [{"source": "/host/x", "target": "/container/x"}],
|
||||||
|
},
|
||||||
|
)
|
||||||
|
assert "/host/x:/container/x" in compose
|
||||||
|
|
||||||
|
|
||||||
|
class TestComputeImageTag:
|
||||||
|
"""Tests for compute_image_tag."""
|
||||||
|
|
||||||
|
def test_is_deterministic(self) -> None:
|
||||||
|
manifest = {"name": "test", "packages": {"apt": ["curl"]}}
|
||||||
|
tag1 = compute_image_tag("My Tool", manifest)
|
||||||
|
tag2 = compute_image_tag("My Tool", manifest)
|
||||||
|
assert tag1 == tag2
|
||||||
|
|
||||||
|
def test_changes_with_content(self) -> None:
|
||||||
|
manifest1 = {"name": "test", "packages": {"apt": ["curl"]}}
|
||||||
|
manifest2 = {"name": "test", "packages": {"apt": ["wget"]}}
|
||||||
|
tag1 = compute_image_tag("test", manifest1)
|
||||||
|
tag2 = compute_image_tag("test", manifest2)
|
||||||
|
assert tag1 != tag2
|
||||||
|
|
||||||
|
def test_lowercases_name(self) -> None:
|
||||||
|
manifest = {"name": "test"}
|
||||||
|
tag = compute_image_tag("My Tool", manifest)
|
||||||
|
assert "my-tool" in tag
|
||||||
|
|
||||||
|
def test_valid_docker_reference(self) -> None:
|
||||||
|
manifest = {"name": "test"}
|
||||||
|
tag = compute_image_tag("test", manifest)
|
||||||
|
assert tag.startswith("headquarter/test-")
|
||||||
|
assert tag.endswith(":latest")
|
||||||
|
|
||||||
|
|
||||||
|
class TestMergeWithConfig:
|
||||||
|
"""Tests for merge_with_config (ConfigProfile only)."""
|
||||||
|
|
||||||
|
def test_no_profile_returns_manifest_unchanged(self) -> None:
|
||||||
|
manifest = {"name": "test"}
|
||||||
|
result = merge_with_config(manifest)
|
||||||
|
assert result["name"] == "test"
|
||||||
|
assert result["_extra_env"] == {}
|
||||||
|
assert result["_extra_volumes"] == []
|
||||||
|
|
||||||
|
def test_profile_env_vars(self) -> None:
|
||||||
|
manifest = {"name": "test"}
|
||||||
|
profile = {"environment_variables": {"FOO": "bar"}}
|
||||||
|
result = merge_with_config(manifest, profile)
|
||||||
|
assert result["_extra_env"]["FOO"] == "bar"
|
||||||
|
|
||||||
|
def test_profile_mounts(self) -> None:
|
||||||
|
manifest = {"name": "test"}
|
||||||
|
profile = {"mounts": [{"source": "/host", "target": "/container"}]}
|
||||||
|
result = merge_with_config(manifest, profile)
|
||||||
|
assert len(result["_extra_volumes"]) == 1
|
||||||
|
|
||||||
|
def test_profile_port_override(self) -> None:
|
||||||
|
manifest = {"name": "test", "default_port": 8080}
|
||||||
|
profile = {"hints": {"port_override": 3000}}
|
||||||
|
result = merge_with_config(manifest, profile)
|
||||||
|
assert result["default_port"] == 3000
|
||||||
|
|
||||||
|
def test_profile_start_command(self) -> None:
|
||||||
|
manifest = {"name": "test", "runtime": {"command": ["/bin/bash"]}}
|
||||||
|
profile = {"hints": {"start_command": "/bin/sh"}}
|
||||||
|
result = merge_with_config(manifest, profile)
|
||||||
|
assert result["runtime"]["command"] == ["/bin/sh"]
|
||||||
|
|
||||||
|
def test_profile_working_directory(self) -> None:
|
||||||
|
manifest = {"name": "test"}
|
||||||
|
profile = {"hints": {"working_directory": "/workspace"}}
|
||||||
|
result = merge_with_config(manifest, profile)
|
||||||
|
assert result["runtime"]["working_dir"] == "/workspace"
|
||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user