bamboo-coding 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- agent/__init__.py +0 -0
- agent/bamboo_git_agent/__init__.py +0 -0
- agent/bamboo_git_agent/capabilities/__init__.py +0 -0
- agent/bamboo_git_agent/capabilities/repo.py +189 -0
- agent/bamboo_git_agent/config.py +76 -0
- agent/bamboo_git_agent/controller_client.py +73 -0
- agent/bamboo_git_agent/journal.py +209 -0
- agent/bamboo_git_agent/main.py +164 -0
- agent/bamboo_git_agent/protocol.py +21 -0
- agent/bamboo_git_agent/status.py +59 -0
- api/__init__.py +2 -0
- api/commit.py +82 -0
- api/repository.py +592 -0
- bamboo_coding/__init__.py +3 -0
- bamboo_coding/client/__init__.py +3 -0
- bamboo_coding/client/capabilities/__init__.py +1 -0
- bamboo_coding/client/capabilities/repo.py +279 -0
- bamboo_coding/client/config.py +179 -0
- bamboo_coding/client/controller_client.py +112 -0
- bamboo_coding/client/journal.py +1 -0
- bamboo_coding/client/main.py +425 -0
- bamboo_coding/client/status.py +59 -0
- bamboo_coding/client/terminal_runtime.py +251 -0
- bamboo_coding/server/__init__.py +3 -0
- bamboo_coding/server/core/__init__.py +1 -0
- bamboo_coding/server/core/config.py +1 -0
- bamboo_coding/server/main.py +39 -0
- bamboo_coding/shared/__init__.py +39 -0
- bamboo_coding/shared/protocol.py +1 -0
- bamboo_coding-0.1.0.dist-info/METADATA +284 -0
- bamboo_coding-0.1.0.dist-info/RECORD +62 -0
- bamboo_coding-0.1.0.dist-info/WHEEL +5 -0
- bamboo_coding-0.1.0.dist-info/entry_points.txt +3 -0
- bamboo_coding-0.1.0.dist-info/top_level.txt +6 -0
- controller/__init__.py +0 -0
- controller/app/__init__.py +0 -0
- controller/app/api/__init__.py +0 -0
- controller/app/api/agents/__init__.py +3 -0
- controller/app/api/agents/ws.py +133 -0
- controller/app/api/public/__init__.py +11 -0
- controller/app/api/public/repos.py +207 -0
- controller/app/api/public/terminals.py +64 -0
- controller/app/core/__init__.py +0 -0
- controller/app/core/config.py +20 -0
- controller/app/core/errors.py +14 -0
- controller/app/db/__init__.py +0 -0
- controller/app/db/models.py +42 -0
- controller/app/db/session.py +74 -0
- controller/app/main.py +71 -0
- controller/app/schemas/__init__.py +23 -0
- controller/app/schemas/agent_messages.py +21 -0
- controller/app/schemas/public.py +3 -0
- controller/app/services/__init__.py +0 -0
- controller/app/services/agents.py +108 -0
- controller/app/services/registrations.py +17 -0
- controller/app/services/repositories.py +85 -0
- controller/app/services/router.py +140 -0
- controller/app/services/tasks.py +104 -0
- controller/app/services/terminals.py +148 -0
- git_utils.py +2238 -0
- shared/__init__.py +39 -0
- shared/protocol.py +253 -0
|
@@ -0,0 +1,23 @@
|
|
|
1
|
+
from .agent_messages import (
|
|
2
|
+
AgentRegisterMessage,
|
|
3
|
+
ClientStatusMessage,
|
|
4
|
+
RegisterRequestMessage,
|
|
5
|
+
TaskEventMessage,
|
|
6
|
+
TaskMessage,
|
|
7
|
+
TaskQueryMessage,
|
|
8
|
+
TaskQueryResultMessage,
|
|
9
|
+
TaskResultMessage,
|
|
10
|
+
)
|
|
11
|
+
from .public import PublicRepositorySummary
|
|
12
|
+
|
|
13
|
+
__all__ = [
|
|
14
|
+
"AgentRegisterMessage",
|
|
15
|
+
"ClientStatusMessage",
|
|
16
|
+
"PublicRepositorySummary",
|
|
17
|
+
"RegisterRequestMessage",
|
|
18
|
+
"TaskEventMessage",
|
|
19
|
+
"TaskMessage",
|
|
20
|
+
"TaskQueryMessage",
|
|
21
|
+
"TaskQueryResultMessage",
|
|
22
|
+
"TaskResultMessage",
|
|
23
|
+
]
|
|
@@ -0,0 +1,21 @@
|
|
|
1
|
+
from shared.protocol import (
|
|
2
|
+
AgentRegisterMessage,
|
|
3
|
+
ClientStatusMessage,
|
|
4
|
+
RegisterRequestMessage,
|
|
5
|
+
TaskEventMessage,
|
|
6
|
+
TaskMessage,
|
|
7
|
+
TaskQueryMessage,
|
|
8
|
+
TaskQueryResultMessage,
|
|
9
|
+
TaskResultMessage,
|
|
10
|
+
)
|
|
11
|
+
|
|
12
|
+
__all__ = [
|
|
13
|
+
"AgentRegisterMessage",
|
|
14
|
+
"ClientStatusMessage",
|
|
15
|
+
"RegisterRequestMessage",
|
|
16
|
+
"TaskEventMessage",
|
|
17
|
+
"TaskMessage",
|
|
18
|
+
"TaskQueryMessage",
|
|
19
|
+
"TaskQueryResultMessage",
|
|
20
|
+
"TaskResultMessage",
|
|
21
|
+
]
|
|
File without changes
|
|
@@ -0,0 +1,108 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
from datetime import UTC, datetime
|
|
4
|
+
import json
|
|
5
|
+
from uuid import uuid4
|
|
6
|
+
|
|
7
|
+
from controller.app.core.errors import AgentApprovalError
|
|
8
|
+
from controller.app.db.session import ControllerDatabase
|
|
9
|
+
|
|
10
|
+
|
|
11
|
+
def _now() -> str:
|
|
12
|
+
return datetime.now(UTC).isoformat()
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
def _decode_agent(row) -> dict | None:
|
|
16
|
+
if row is None:
|
|
17
|
+
return None
|
|
18
|
+
payload = dict(row)
|
|
19
|
+
payload["client_info"] = json.loads(payload.pop("client_info_json"))
|
|
20
|
+
return payload
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
class AgentService:
|
|
24
|
+
def __init__(self, database: ControllerDatabase):
|
|
25
|
+
self.database = database
|
|
26
|
+
|
|
27
|
+
def create_registration_request(self, client_id: str) -> dict:
|
|
28
|
+
now = _now()
|
|
29
|
+
with self.database.connection() as conn:
|
|
30
|
+
row = conn.execute("SELECT * FROM agents WHERE client_id = ?", (client_id,)).fetchone()
|
|
31
|
+
if row is None:
|
|
32
|
+
agent_id = str(uuid4())
|
|
33
|
+
conn.execute(
|
|
34
|
+
"INSERT INTO agents (agent_id, client_id, token, approval_status, status, client_info_json, created_at, updated_at, last_seen_at) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)",
|
|
35
|
+
(agent_id, client_id, None, "pending", "pending", json.dumps({}), now, now, None),
|
|
36
|
+
)
|
|
37
|
+
row = conn.execute("SELECT * FROM agents WHERE client_id = ?", (client_id,)).fetchone()
|
|
38
|
+
return _decode_agent(row)
|
|
39
|
+
|
|
40
|
+
def approve_registration(self, client_id: str, token: str) -> dict:
|
|
41
|
+
now = _now()
|
|
42
|
+
with self.database.connection() as conn:
|
|
43
|
+
row = conn.execute("SELECT * FROM agents WHERE client_id = ?", (client_id,)).fetchone()
|
|
44
|
+
if row is None:
|
|
45
|
+
self.create_registration_request(client_id)
|
|
46
|
+
conn.execute(
|
|
47
|
+
"UPDATE agents SET token = ?, approval_status = ?, updated_at = ? WHERE client_id = ?",
|
|
48
|
+
(token, "approved", now, client_id),
|
|
49
|
+
)
|
|
50
|
+
row = conn.execute("SELECT * FROM agents WHERE client_id = ?", (client_id,)).fetchone()
|
|
51
|
+
return _decode_agent(row)
|
|
52
|
+
|
|
53
|
+
def approve_agent(self, agent_id: str, token: str) -> dict:
|
|
54
|
+
now = _now()
|
|
55
|
+
with self.database.connection() as conn:
|
|
56
|
+
row = conn.execute("SELECT * FROM agents WHERE agent_id = ?", (agent_id,)).fetchone()
|
|
57
|
+
if row is None:
|
|
58
|
+
raise AgentApprovalError(f"Unknown agent: {agent_id}")
|
|
59
|
+
conn.execute(
|
|
60
|
+
"UPDATE agents SET token = ?, approval_status = ?, updated_at = ? WHERE agent_id = ?",
|
|
61
|
+
(token, "approved", now, agent_id),
|
|
62
|
+
)
|
|
63
|
+
row = conn.execute("SELECT * FROM agents WHERE agent_id = ?", (agent_id,)).fetchone()
|
|
64
|
+
return _decode_agent(row)
|
|
65
|
+
|
|
66
|
+
def register_connection(self, client_id: str, token: str, client_info: dict) -> dict:
|
|
67
|
+
now = _now()
|
|
68
|
+
with self.database.connection() as conn:
|
|
69
|
+
row = conn.execute(
|
|
70
|
+
"SELECT * FROM agents WHERE client_id = ? AND token = ?",
|
|
71
|
+
(client_id, token),
|
|
72
|
+
).fetchone()
|
|
73
|
+
if row is None or row["approval_status"] != "approved":
|
|
74
|
+
raise AgentApprovalError(f"Agent {client_id} is not approved")
|
|
75
|
+
conn.execute(
|
|
76
|
+
"UPDATE agents SET status = ?, client_info_json = ?, updated_at = ?, last_seen_at = ? WHERE client_id = ?",
|
|
77
|
+
("connected", json.dumps(client_info), now, now, client_id),
|
|
78
|
+
)
|
|
79
|
+
row = conn.execute("SELECT * FROM agents WHERE client_id = ?", (client_id,)).fetchone()
|
|
80
|
+
return _decode_agent(row)
|
|
81
|
+
|
|
82
|
+
def get_agent_by_token(self, token: str) -> dict | None:
|
|
83
|
+
with self.database.connection() as conn:
|
|
84
|
+
row = conn.execute("SELECT * FROM agents WHERE token = ?", (token,)).fetchone()
|
|
85
|
+
return _decode_agent(row)
|
|
86
|
+
|
|
87
|
+
def get_agent_by_client_id(self, client_id: str) -> dict | None:
|
|
88
|
+
with self.database.connection() as conn:
|
|
89
|
+
row = conn.execute("SELECT * FROM agents WHERE client_id = ?", (client_id,)).fetchone()
|
|
90
|
+
return _decode_agent(row)
|
|
91
|
+
|
|
92
|
+
def get_agent_by_id(self, agent_id: str) -> dict | None:
|
|
93
|
+
with self.database.connection() as conn:
|
|
94
|
+
row = conn.execute("SELECT * FROM agents WHERE agent_id = ?", (agent_id,)).fetchone()
|
|
95
|
+
return _decode_agent(row)
|
|
96
|
+
|
|
97
|
+
def list_agents(self) -> list[dict]:
|
|
98
|
+
with self.database.connection() as conn:
|
|
99
|
+
rows = conn.execute("SELECT * FROM agents ORDER BY created_at ASC").fetchall()
|
|
100
|
+
return [_decode_agent(row) for row in rows]
|
|
101
|
+
|
|
102
|
+
def mark_disconnected(self, agent_id: str) -> None:
|
|
103
|
+
now = _now()
|
|
104
|
+
with self.database.connection() as conn:
|
|
105
|
+
conn.execute(
|
|
106
|
+
"UPDATE agents SET status = ?, updated_at = ? WHERE agent_id = ?",
|
|
107
|
+
("disconnected", now, agent_id),
|
|
108
|
+
)
|
|
@@ -0,0 +1,17 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
from controller.app.services.agents import AgentService
|
|
4
|
+
from shared.protocol import RegisterRequestMessage
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
class RegistrationService:
|
|
8
|
+
def __init__(self, agents: AgentService):
|
|
9
|
+
self.agents = agents
|
|
10
|
+
|
|
11
|
+
def ensure_pending(self, client_id: str) -> RegisterRequestMessage:
|
|
12
|
+
self.agents.create_registration_request(client_id)
|
|
13
|
+
return RegisterRequestMessage(approved=False, client_id=client_id, message="pending approval")
|
|
14
|
+
|
|
15
|
+
def approve(self, client_id: str, token: str) -> RegisterRequestMessage:
|
|
16
|
+
self.agents.approve_registration(client_id=client_id, token=token)
|
|
17
|
+
return RegisterRequestMessage(approved=True, client_id=client_id, token=token, message="approved")
|
|
@@ -0,0 +1,85 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
from datetime import UTC, datetime
|
|
4
|
+
from pathlib import Path
|
|
5
|
+
import os
|
|
6
|
+
|
|
7
|
+
from controller.app.core.errors import RepositoryNotFoundError
|
|
8
|
+
from controller.app.db.session import ControllerDatabase
|
|
9
|
+
from shared.protocol import PublicRepositorySummary, PublicRootSummary
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
def _now() -> str:
|
|
13
|
+
return datetime.now(UTC).isoformat()
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
def _normalize_path(path_value: str) -> str:
|
|
17
|
+
return str(Path(path_value).expanduser().resolve())
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
class RepositoryService:
|
|
21
|
+
def __init__(self, database: ControllerDatabase):
|
|
22
|
+
self.database = database
|
|
23
|
+
|
|
24
|
+
def upsert_repository_binding(self, agent_id: str, repo_id: str, name: str, path: str, status: str) -> dict:
|
|
25
|
+
now = _now()
|
|
26
|
+
normalized_path = _normalize_path(path)
|
|
27
|
+
storage_key = _normalize_path(repo_id) if repo_id else normalized_path
|
|
28
|
+
with self.database.connection() as conn:
|
|
29
|
+
conn.execute(
|
|
30
|
+
"""
|
|
31
|
+
INSERT INTO repositories (repo_id, agent_id, name, path, status, created_at, updated_at)
|
|
32
|
+
VALUES (?, ?, ?, ?, ?, ?, ?)
|
|
33
|
+
ON CONFLICT(repo_id) DO UPDATE SET
|
|
34
|
+
agent_id = excluded.agent_id,
|
|
35
|
+
name = excluded.name,
|
|
36
|
+
path = excluded.path,
|
|
37
|
+
status = excluded.status,
|
|
38
|
+
updated_at = excluded.updated_at
|
|
39
|
+
""",
|
|
40
|
+
(storage_key, agent_id, name, normalized_path, status, now, now),
|
|
41
|
+
)
|
|
42
|
+
return self.get_repository_binding(storage_key)
|
|
43
|
+
|
|
44
|
+
def upsert_root_binding(self, agent_id: str, path: str, name: str, status: str) -> dict:
|
|
45
|
+
normalized_path = _normalize_path(path)
|
|
46
|
+
return self.upsert_repository_binding(agent_id=agent_id, repo_id=normalized_path, name=name, path=normalized_path, status=status)
|
|
47
|
+
|
|
48
|
+
def list_public_repositories(self) -> list[PublicRepositorySummary]:
|
|
49
|
+
return self.list_public_roots()
|
|
50
|
+
|
|
51
|
+
def list_public_roots(self) -> list[PublicRootSummary]:
|
|
52
|
+
with self.database.connection() as conn:
|
|
53
|
+
rows = conn.execute("SELECT path, name, status FROM repositories ORDER BY path ASC").fetchall()
|
|
54
|
+
return [PublicRootSummary(path=row["path"], name=row["name"], status=row["status"]) for row in rows]
|
|
55
|
+
|
|
56
|
+
def get_repository_binding(self, repo_id: str) -> dict | None:
|
|
57
|
+
with self.database.connection() as conn:
|
|
58
|
+
row = conn.execute("SELECT * FROM repositories WHERE repo_id = ?", (_normalize_path(repo_id),)).fetchone()
|
|
59
|
+
return dict(row) if row is not None else None
|
|
60
|
+
|
|
61
|
+
def get_root_binding(self, path: str) -> dict | None:
|
|
62
|
+
return self.get_repository_binding(path)
|
|
63
|
+
|
|
64
|
+
def require_repository_binding(self, repo_id: str) -> dict:
|
|
65
|
+
binding = self.get_repository_binding(repo_id)
|
|
66
|
+
if binding is None:
|
|
67
|
+
raise RepositoryNotFoundError(f"Unknown repository: {repo_id}")
|
|
68
|
+
return binding
|
|
69
|
+
|
|
70
|
+
def find_root_bindings_for_path(self, path: str) -> list[dict]:
|
|
71
|
+
normalized_target = _normalize_path(path)
|
|
72
|
+
with self.database.connection() as conn:
|
|
73
|
+
rows = conn.execute("SELECT * FROM repositories ORDER BY LENGTH(path) DESC, path ASC").fetchall()
|
|
74
|
+
bindings: list[dict] = []
|
|
75
|
+
for row in rows:
|
|
76
|
+
binding = dict(row)
|
|
77
|
+
binding_path = _normalize_path(binding["path"])
|
|
78
|
+
try:
|
|
79
|
+
contains = os.path.commonpath([binding_path, normalized_target]) == binding_path
|
|
80
|
+
except ValueError:
|
|
81
|
+
contains = False
|
|
82
|
+
if contains:
|
|
83
|
+
binding["path"] = binding_path
|
|
84
|
+
bindings.append(binding)
|
|
85
|
+
return bindings
|
|
@@ -0,0 +1,140 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import asyncio
|
|
4
|
+
from dataclasses import dataclass
|
|
5
|
+
import os
|
|
6
|
+
from pathlib import Path
|
|
7
|
+
from typing import Awaitable, Callable
|
|
8
|
+
from uuid import uuid4
|
|
9
|
+
|
|
10
|
+
from controller.app.services.repositories import RepositoryService
|
|
11
|
+
from controller.app.services.tasks import TaskService
|
|
12
|
+
from shared.protocol import TaskMessage
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
SendJson = Callable[[dict], Awaitable[None]]
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
def _normalize_path(path_value: str) -> str:
|
|
19
|
+
return str(Path(path_value).expanduser().resolve())
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
def _path_is_within_root(path_value: str, root_path: str) -> bool:
|
|
23
|
+
if not root_path:
|
|
24
|
+
return False
|
|
25
|
+
normalized_target = _normalize_path(path_value)
|
|
26
|
+
normalized_root = _normalize_path(root_path)
|
|
27
|
+
try:
|
|
28
|
+
return os.path.commonpath([normalized_root, normalized_target]) == normalized_root
|
|
29
|
+
except ValueError:
|
|
30
|
+
return False
|
|
31
|
+
|
|
32
|
+
|
|
33
|
+
@dataclass
|
|
34
|
+
class LiveAgent:
|
|
35
|
+
agent_id: str
|
|
36
|
+
send_json: SendJson
|
|
37
|
+
capabilities: list[dict]
|
|
38
|
+
roots: list[dict]
|
|
39
|
+
|
|
40
|
+
def supports(self, capability: str, action: str) -> bool:
|
|
41
|
+
for item in self.capabilities:
|
|
42
|
+
if item.get("capability") != capability:
|
|
43
|
+
continue
|
|
44
|
+
if action in (item.get("actions") or []):
|
|
45
|
+
return True
|
|
46
|
+
return False
|
|
47
|
+
|
|
48
|
+
|
|
49
|
+
class TaskRouter:
|
|
50
|
+
def __init__(self, repositories: RepositoryService, tasks: TaskService, timeout_seconds: float = 15.0):
|
|
51
|
+
self.repositories = repositories
|
|
52
|
+
self.tasks = tasks
|
|
53
|
+
self.timeout_seconds = timeout_seconds
|
|
54
|
+
self.live_agents: dict[str, LiveAgent] = {}
|
|
55
|
+
self.pending_tasks: dict[str, asyncio.Future] = {}
|
|
56
|
+
self.pending_tasks_history: list[str] = []
|
|
57
|
+
|
|
58
|
+
def attach_agent(
|
|
59
|
+
self,
|
|
60
|
+
*,
|
|
61
|
+
agent_id: str,
|
|
62
|
+
send_json: SendJson,
|
|
63
|
+
capabilities: list[dict],
|
|
64
|
+
roots: list[dict] | None = None,
|
|
65
|
+
repositories: list[dict] | None = None,
|
|
66
|
+
) -> None:
|
|
67
|
+
self.live_agents[agent_id] = LiveAgent(
|
|
68
|
+
agent_id=agent_id,
|
|
69
|
+
send_json=send_json,
|
|
70
|
+
capabilities=capabilities,
|
|
71
|
+
roots=roots if roots is not None else (repositories or []),
|
|
72
|
+
)
|
|
73
|
+
|
|
74
|
+
def detach_agent(self, agent_id: str) -> None:
|
|
75
|
+
self.live_agents.pop(agent_id, None)
|
|
76
|
+
|
|
77
|
+
async def dispatch_task(self, path: str, capability: str, action: str, params: dict, agent_id: str | None = None) -> dict:
|
|
78
|
+
normalized_path = _normalize_path(path)
|
|
79
|
+
selected_agent_id = agent_id
|
|
80
|
+
|
|
81
|
+
if selected_agent_id is not None:
|
|
82
|
+
live_agent = self.live_agents.get(selected_agent_id)
|
|
83
|
+
if live_agent is None:
|
|
84
|
+
raise RuntimeError(f"No live agent for selected client {selected_agent_id}")
|
|
85
|
+
if not live_agent.supports(capability, action):
|
|
86
|
+
raise RuntimeError(f"Selected client {selected_agent_id} does not support {capability}:{action}")
|
|
87
|
+
if not any(_path_is_within_root(normalized_path, root.get("path", "")) for root in live_agent.roots):
|
|
88
|
+
raise RuntimeError(f"Path {normalized_path} is outside selected client roots")
|
|
89
|
+
else:
|
|
90
|
+
candidate_bindings = self.repositories.find_root_bindings_for_path(normalized_path)
|
|
91
|
+
routed: list[tuple[dict, LiveAgent]] = []
|
|
92
|
+
for binding in candidate_bindings:
|
|
93
|
+
live_agent = self.live_agents.get(binding["agent_id"])
|
|
94
|
+
if live_agent is None:
|
|
95
|
+
continue
|
|
96
|
+
if not live_agent.supports(capability, action):
|
|
97
|
+
continue
|
|
98
|
+
routed.append((binding, live_agent))
|
|
99
|
+
|
|
100
|
+
if not routed:
|
|
101
|
+
raise RuntimeError(f"No live agent for path {normalized_path}")
|
|
102
|
+
if len(routed) > 1:
|
|
103
|
+
raise RuntimeError(f"Ambiguous path owner for {normalized_path}")
|
|
104
|
+
|
|
105
|
+
binding, live_agent = routed[0]
|
|
106
|
+
selected_agent_id = binding["agent_id"]
|
|
107
|
+
|
|
108
|
+
task_id = str(uuid4())
|
|
109
|
+
task = TaskMessage(task_id=task_id, path=normalized_path, capability=capability, action=action, params=params)
|
|
110
|
+
self.tasks.create_task(
|
|
111
|
+
task_id=task_id,
|
|
112
|
+
repo_id=normalized_path,
|
|
113
|
+
capability=capability,
|
|
114
|
+
action=action,
|
|
115
|
+
params=params,
|
|
116
|
+
agent_id=selected_agent_id,
|
|
117
|
+
)
|
|
118
|
+
loop = asyncio.get_running_loop()
|
|
119
|
+
future = loop.create_future()
|
|
120
|
+
self.pending_tasks[task_id] = future
|
|
121
|
+
self.pending_tasks_history.append(task_id)
|
|
122
|
+
await live_agent.send_json(task.model_dump())
|
|
123
|
+
try:
|
|
124
|
+
return await asyncio.wait_for(future, timeout=self.timeout_seconds)
|
|
125
|
+
finally:
|
|
126
|
+
self.pending_tasks.pop(task_id, None)
|
|
127
|
+
|
|
128
|
+
async def handle_event(self, *, task_id: str, status: str, message: str | None = None, details: dict | None = None) -> dict:
|
|
129
|
+
return self.tasks.append_event(task_id, status=status, message=message, details=details)
|
|
130
|
+
|
|
131
|
+
async def handle_result(self, task_id: str, success: bool, result: dict | None, error: str | None) -> dict:
|
|
132
|
+
stored = self.tasks.store_result(task_id, success=success, result=result, error=error)
|
|
133
|
+
future = self.pending_tasks.get(task_id)
|
|
134
|
+
if future is None or future.done() or not stored["inserted"]:
|
|
135
|
+
return stored
|
|
136
|
+
if success:
|
|
137
|
+
future.set_result(result or {})
|
|
138
|
+
else:
|
|
139
|
+
future.set_exception(RuntimeError(error or "Task failed"))
|
|
140
|
+
return stored
|
|
@@ -0,0 +1,104 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
from datetime import UTC, datetime
|
|
4
|
+
import json
|
|
5
|
+
|
|
6
|
+
from controller.app.core.errors import TaskNotFoundError
|
|
7
|
+
from controller.app.db.session import ControllerDatabase
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
def _now() -> str:
|
|
11
|
+
return datetime.now(UTC).isoformat()
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
def _decode_task(row) -> dict | None:
|
|
15
|
+
if row is None:
|
|
16
|
+
return None
|
|
17
|
+
payload = dict(row)
|
|
18
|
+
payload["params"] = json.loads(payload.pop("params_json"))
|
|
19
|
+
result_json = payload.pop("result_json")
|
|
20
|
+
payload["result"] = json.loads(result_json) if result_json else None
|
|
21
|
+
payload["success"] = None if payload["result_success"] is None else bool(payload["result_success"])
|
|
22
|
+
return payload
|
|
23
|
+
|
|
24
|
+
|
|
25
|
+
class TaskService:
|
|
26
|
+
def __init__(self, database: ControllerDatabase):
|
|
27
|
+
self.database = database
|
|
28
|
+
|
|
29
|
+
def create_task(self, task_id: str, repo_id: str, capability: str, action: str, params: dict, agent_id: str | None = None) -> dict:
|
|
30
|
+
now = _now()
|
|
31
|
+
with self.database.connection() as conn:
|
|
32
|
+
conn.execute(
|
|
33
|
+
"INSERT INTO tasks (task_id, repo_id, agent_id, capability, action, params_json, status, created_at, updated_at, terminal_at, result_success, result_json, error) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)",
|
|
34
|
+
(task_id, repo_id, agent_id, capability, action, json.dumps(params), "queued", now, now, None, None, None, None),
|
|
35
|
+
)
|
|
36
|
+
row = conn.execute("SELECT * FROM tasks WHERE task_id = ?", (task_id,)).fetchone()
|
|
37
|
+
return _decode_task(row)
|
|
38
|
+
|
|
39
|
+
def append_event(self, task_id: str, status: str, message: str | None = None, details: dict | None = None) -> dict:
|
|
40
|
+
now = _now()
|
|
41
|
+
details = details or {}
|
|
42
|
+
with self.database.connection() as conn:
|
|
43
|
+
task = conn.execute("SELECT * FROM tasks WHERE task_id = ?", (task_id,)).fetchone()
|
|
44
|
+
if task is None:
|
|
45
|
+
raise TaskNotFoundError(f"Unknown task: {task_id}")
|
|
46
|
+
conn.execute(
|
|
47
|
+
"INSERT INTO task_events (task_id, status, message, details_json, created_at) VALUES (?, ?, ?, ?, ?)",
|
|
48
|
+
(task_id, status, message, json.dumps(details), now),
|
|
49
|
+
)
|
|
50
|
+
conn.execute(
|
|
51
|
+
"UPDATE tasks SET status = ?, updated_at = ? WHERE task_id = ?",
|
|
52
|
+
(status, now, task_id),
|
|
53
|
+
)
|
|
54
|
+
return {"task_id": task_id, "status": status, "message": message, "details": details}
|
|
55
|
+
|
|
56
|
+
def list_events(self, task_id: str) -> list[dict]:
|
|
57
|
+
with self.database.connection() as conn:
|
|
58
|
+
rows = conn.execute(
|
|
59
|
+
"SELECT task_id, status, message, details_json, created_at FROM task_events WHERE task_id = ? ORDER BY id ASC",
|
|
60
|
+
(task_id,),
|
|
61
|
+
).fetchall()
|
|
62
|
+
return [
|
|
63
|
+
{
|
|
64
|
+
"task_id": row["task_id"],
|
|
65
|
+
"status": row["status"],
|
|
66
|
+
"message": row["message"],
|
|
67
|
+
"details": json.loads(row["details_json"]),
|
|
68
|
+
"created_at": row["created_at"],
|
|
69
|
+
}
|
|
70
|
+
for row in rows
|
|
71
|
+
]
|
|
72
|
+
|
|
73
|
+
def store_result(self, task_id: str, success: bool, result: dict | None = None, error: str | None = None) -> dict:
|
|
74
|
+
now = _now()
|
|
75
|
+
result = result or {}
|
|
76
|
+
terminal_status = "succeeded" if success else "failed"
|
|
77
|
+
with self.database.connection() as conn:
|
|
78
|
+
row = conn.execute("SELECT * FROM tasks WHERE task_id = ?", (task_id,)).fetchone()
|
|
79
|
+
if row is None:
|
|
80
|
+
raise TaskNotFoundError(f"Unknown task: {task_id}")
|
|
81
|
+
if row["result_success"] is not None:
|
|
82
|
+
return {"inserted": False, "task": _decode_task(row)}
|
|
83
|
+
conn.execute(
|
|
84
|
+
"UPDATE tasks SET status = ?, updated_at = ?, terminal_at = ?, result_success = ?, result_json = ?, error = ? WHERE task_id = ?",
|
|
85
|
+
(terminal_status, now, now, int(success), json.dumps(result), error, task_id),
|
|
86
|
+
)
|
|
87
|
+
updated = conn.execute("SELECT * FROM tasks WHERE task_id = ?", (task_id,)).fetchone()
|
|
88
|
+
return {"inserted": True, "task": _decode_task(updated)}
|
|
89
|
+
|
|
90
|
+
def get_task(self, task_id: str) -> dict | None:
|
|
91
|
+
with self.database.connection() as conn:
|
|
92
|
+
row = conn.execute("SELECT * FROM tasks WHERE task_id = ?", (task_id,)).fetchone()
|
|
93
|
+
return _decode_task(row)
|
|
94
|
+
|
|
95
|
+
def get_result(self, task_id: str) -> dict | None:
|
|
96
|
+
task = self.get_task(task_id)
|
|
97
|
+
if task is None or task["success"] is None:
|
|
98
|
+
return None
|
|
99
|
+
return {
|
|
100
|
+
"task_id": task_id,
|
|
101
|
+
"success": task["success"],
|
|
102
|
+
"result": task["result"] or {},
|
|
103
|
+
"error": task["error"],
|
|
104
|
+
}
|
|
@@ -0,0 +1,148 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
from dataclasses import dataclass
|
|
4
|
+
from typing import Awaitable, Callable
|
|
5
|
+
|
|
6
|
+
from controller.app.services.router import TaskRouter, _normalize_path, _path_is_within_root
|
|
7
|
+
from shared.protocol import (
|
|
8
|
+
TerminalCloseMessage,
|
|
9
|
+
TerminalClosedMessage,
|
|
10
|
+
TerminalHeartbeatMessage,
|
|
11
|
+
TerminalInputMessage,
|
|
12
|
+
TerminalOpenMessage,
|
|
13
|
+
TerminalOutputMessage,
|
|
14
|
+
TerminalResizeMessage,
|
|
15
|
+
TerminalStatusMessage,
|
|
16
|
+
)
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
SendJson = Callable[[dict], Awaitable[None]]
|
|
20
|
+
|
|
21
|
+
|
|
22
|
+
@dataclass
|
|
23
|
+
class BrowserTerminalSession:
|
|
24
|
+
terminal_id: str
|
|
25
|
+
agent_id: str
|
|
26
|
+
path: str
|
|
27
|
+
cols: int
|
|
28
|
+
rows: int
|
|
29
|
+
send_json: SendJson
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
class TerminalRelayService:
|
|
33
|
+
def __init__(self, router: TaskRouter):
|
|
34
|
+
self.router = router
|
|
35
|
+
self.sessions: dict[str, BrowserTerminalSession] = {}
|
|
36
|
+
|
|
37
|
+
async def open_browser_terminal(
|
|
38
|
+
self,
|
|
39
|
+
*,
|
|
40
|
+
terminal_id: str,
|
|
41
|
+
agent_id: str,
|
|
42
|
+
path: str,
|
|
43
|
+
cols: int,
|
|
44
|
+
rows: int,
|
|
45
|
+
send_json: SendJson,
|
|
46
|
+
) -> str:
|
|
47
|
+
normalized_path = _normalize_path(path)
|
|
48
|
+
live_agent = self.router.live_agents.get(agent_id)
|
|
49
|
+
if live_agent is None:
|
|
50
|
+
raise RuntimeError(f"No live agent for selected client {agent_id}")
|
|
51
|
+
if not live_agent.supports('terminal.open', 'post'):
|
|
52
|
+
raise RuntimeError(f"Selected client {agent_id} does not support terminal.open:post")
|
|
53
|
+
if not any(_path_is_within_root(normalized_path, root.get('path', '')) for root in live_agent.roots):
|
|
54
|
+
raise RuntimeError(f"Path {normalized_path} is outside selected client roots")
|
|
55
|
+
|
|
56
|
+
session = BrowserTerminalSession(
|
|
57
|
+
terminal_id=terminal_id,
|
|
58
|
+
agent_id=agent_id,
|
|
59
|
+
path=normalized_path,
|
|
60
|
+
cols=cols,
|
|
61
|
+
rows=rows,
|
|
62
|
+
send_json=send_json,
|
|
63
|
+
)
|
|
64
|
+
self.sessions[terminal_id] = session
|
|
65
|
+
await send_json(TerminalStatusMessage(terminal_id=terminal_id, status='connecting').model_dump())
|
|
66
|
+
await live_agent.send_json(
|
|
67
|
+
TerminalOpenMessage(terminal_id=terminal_id, path=normalized_path, cols=cols, rows=rows).model_dump()
|
|
68
|
+
)
|
|
69
|
+
return terminal_id
|
|
70
|
+
|
|
71
|
+
async def forward_browser_message(self, terminal_id: str, payload: dict) -> None:
|
|
72
|
+
session = self.sessions.get(terminal_id)
|
|
73
|
+
if session is None:
|
|
74
|
+
return
|
|
75
|
+
live_agent = self.router.live_agents.get(session.agent_id)
|
|
76
|
+
if live_agent is None:
|
|
77
|
+
await self.handle_terminal_closed(TerminalClosedMessage(terminal_id=terminal_id, reason='client_disconnected'))
|
|
78
|
+
return
|
|
79
|
+
|
|
80
|
+
message_type = payload.get('type')
|
|
81
|
+
if message_type == 'terminal_input':
|
|
82
|
+
message = TerminalInputMessage(**payload)
|
|
83
|
+
await live_agent.send_json(message.model_dump())
|
|
84
|
+
elif message_type == 'terminal_resize':
|
|
85
|
+
message = TerminalResizeMessage(**payload)
|
|
86
|
+
await live_agent.send_json(message.model_dump())
|
|
87
|
+
elif message_type == 'terminal_close':
|
|
88
|
+
message = TerminalCloseMessage(**payload)
|
|
89
|
+
await live_agent.send_json(message.model_dump())
|
|
90
|
+
|
|
91
|
+
async def handle_terminal_status(self, message: TerminalStatusMessage) -> None:
|
|
92
|
+
await self._safe_send_to_browser(message.terminal_id, message.model_dump())
|
|
93
|
+
|
|
94
|
+
async def handle_terminal_output(self, message: TerminalOutputMessage) -> None:
|
|
95
|
+
await self._safe_send_to_browser(message.terminal_id, message.model_dump())
|
|
96
|
+
|
|
97
|
+
async def handle_terminal_heartbeat(self, message: TerminalHeartbeatMessage) -> None:
|
|
98
|
+
await self._safe_send_to_browser(message.terminal_id, message.model_dump())
|
|
99
|
+
|
|
100
|
+
async def handle_terminal_closed(self, message: TerminalClosedMessage) -> None:
|
|
101
|
+
session = self.sessions.pop(message.terminal_id, None)
|
|
102
|
+
if session is None:
|
|
103
|
+
return
|
|
104
|
+
try:
|
|
105
|
+
await session.send_json(message.model_dump())
|
|
106
|
+
except Exception:
|
|
107
|
+
return
|
|
108
|
+
|
|
109
|
+
async def close_browser_terminal(self, terminal_id: str, reason: str = 'browser_disconnected') -> None:
|
|
110
|
+
session = self.sessions.pop(terminal_id, None)
|
|
111
|
+
if session is None:
|
|
112
|
+
return
|
|
113
|
+
live_agent = self.router.live_agents.get(session.agent_id)
|
|
114
|
+
if live_agent is None:
|
|
115
|
+
return
|
|
116
|
+
if live_agent.supports('terminal.close', 'post'):
|
|
117
|
+
await live_agent.send_json(TerminalCloseMessage(terminal_id=terminal_id, reason=reason).model_dump())
|
|
118
|
+
|
|
119
|
+
async def close_agent_sessions(self, agent_id: str, reason: str = 'client_disconnected') -> None:
|
|
120
|
+
terminal_ids = [terminal_id for terminal_id, session in self.sessions.items() if session.agent_id == agent_id]
|
|
121
|
+
for terminal_id in terminal_ids:
|
|
122
|
+
session = self.sessions.pop(terminal_id, None)
|
|
123
|
+
if session is None:
|
|
124
|
+
continue
|
|
125
|
+
try:
|
|
126
|
+
await session.send_json(TerminalClosedMessage(terminal_id=terminal_id, reason=reason).model_dump())
|
|
127
|
+
except Exception:
|
|
128
|
+
continue
|
|
129
|
+
|
|
130
|
+
async def _safe_send_to_browser(self, terminal_id: str, payload: dict) -> None:
|
|
131
|
+
session = self.sessions.get(terminal_id)
|
|
132
|
+
if session is None:
|
|
133
|
+
return
|
|
134
|
+
try:
|
|
135
|
+
await session.send_json(payload)
|
|
136
|
+
except Exception:
|
|
137
|
+
self.sessions.pop(terminal_id, None)
|
|
138
|
+
await self.close_remote_terminal(terminal_id, session.agent_id, reason='browser_send_failed')
|
|
139
|
+
|
|
140
|
+
async def close_remote_terminal(self, terminal_id: str, agent_id: str, reason: str) -> None:
|
|
141
|
+
live_agent = self.router.live_agents.get(agent_id)
|
|
142
|
+
if live_agent is None:
|
|
143
|
+
return
|
|
144
|
+
if live_agent.supports('terminal.close', 'post'):
|
|
145
|
+
try:
|
|
146
|
+
await live_agent.send_json(TerminalCloseMessage(terminal_id=terminal_id, reason=reason).model_dump())
|
|
147
|
+
except Exception:
|
|
148
|
+
return
|