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.
Files changed (62) hide show
  1. agent/__init__.py +0 -0
  2. agent/bamboo_git_agent/__init__.py +0 -0
  3. agent/bamboo_git_agent/capabilities/__init__.py +0 -0
  4. agent/bamboo_git_agent/capabilities/repo.py +189 -0
  5. agent/bamboo_git_agent/config.py +76 -0
  6. agent/bamboo_git_agent/controller_client.py +73 -0
  7. agent/bamboo_git_agent/journal.py +209 -0
  8. agent/bamboo_git_agent/main.py +164 -0
  9. agent/bamboo_git_agent/protocol.py +21 -0
  10. agent/bamboo_git_agent/status.py +59 -0
  11. api/__init__.py +2 -0
  12. api/commit.py +82 -0
  13. api/repository.py +592 -0
  14. bamboo_coding/__init__.py +3 -0
  15. bamboo_coding/client/__init__.py +3 -0
  16. bamboo_coding/client/capabilities/__init__.py +1 -0
  17. bamboo_coding/client/capabilities/repo.py +279 -0
  18. bamboo_coding/client/config.py +179 -0
  19. bamboo_coding/client/controller_client.py +112 -0
  20. bamboo_coding/client/journal.py +1 -0
  21. bamboo_coding/client/main.py +425 -0
  22. bamboo_coding/client/status.py +59 -0
  23. bamboo_coding/client/terminal_runtime.py +251 -0
  24. bamboo_coding/server/__init__.py +3 -0
  25. bamboo_coding/server/core/__init__.py +1 -0
  26. bamboo_coding/server/core/config.py +1 -0
  27. bamboo_coding/server/main.py +39 -0
  28. bamboo_coding/shared/__init__.py +39 -0
  29. bamboo_coding/shared/protocol.py +1 -0
  30. bamboo_coding-0.1.0.dist-info/METADATA +284 -0
  31. bamboo_coding-0.1.0.dist-info/RECORD +62 -0
  32. bamboo_coding-0.1.0.dist-info/WHEEL +5 -0
  33. bamboo_coding-0.1.0.dist-info/entry_points.txt +3 -0
  34. bamboo_coding-0.1.0.dist-info/top_level.txt +6 -0
  35. controller/__init__.py +0 -0
  36. controller/app/__init__.py +0 -0
  37. controller/app/api/__init__.py +0 -0
  38. controller/app/api/agents/__init__.py +3 -0
  39. controller/app/api/agents/ws.py +133 -0
  40. controller/app/api/public/__init__.py +11 -0
  41. controller/app/api/public/repos.py +207 -0
  42. controller/app/api/public/terminals.py +64 -0
  43. controller/app/core/__init__.py +0 -0
  44. controller/app/core/config.py +20 -0
  45. controller/app/core/errors.py +14 -0
  46. controller/app/db/__init__.py +0 -0
  47. controller/app/db/models.py +42 -0
  48. controller/app/db/session.py +74 -0
  49. controller/app/main.py +71 -0
  50. controller/app/schemas/__init__.py +23 -0
  51. controller/app/schemas/agent_messages.py +21 -0
  52. controller/app/schemas/public.py +3 -0
  53. controller/app/services/__init__.py +0 -0
  54. controller/app/services/agents.py +108 -0
  55. controller/app/services/registrations.py +17 -0
  56. controller/app/services/repositories.py +85 -0
  57. controller/app/services/router.py +140 -0
  58. controller/app/services/tasks.py +104 -0
  59. controller/app/services/terminals.py +148 -0
  60. git_utils.py +2238 -0
  61. shared/__init__.py +39 -0
  62. shared/protocol.py +253 -0
@@ -0,0 +1,279 @@
1
+ from __future__ import annotations
2
+
3
+ from pathlib import Path
4
+ from typing import Callable
5
+
6
+ import git_utils
7
+ from bamboo_coding.shared.protocol import TaskMessage
8
+
9
+
10
+ def _allowed_roots(repositories: dict[str, str] | list[str] | tuple[str, ...]) -> list[str]:
11
+ if isinstance(repositories, dict):
12
+ values = list(repositories.values()) or list(repositories.keys())
13
+ else:
14
+ values = list(repositories)
15
+ return [str(Path(item).expanduser().resolve()) for item in values]
16
+
17
+
18
+ def _resolve_allowed_path(task: TaskMessage, repositories: dict[str, str] | list[str] | tuple[str, ...]) -> str:
19
+ return _ensure_allowed_path(task.path, repositories)
20
+
21
+
22
+ def _ensure_allowed_path(raw_path: str, repositories: dict[str, str] | list[str] | tuple[str, ...]) -> str:
23
+ target = str(Path(raw_path).expanduser().resolve())
24
+ for root in _allowed_roots(repositories):
25
+ try:
26
+ Path(target).relative_to(root)
27
+ return target
28
+ except ValueError:
29
+ continue
30
+ raise PermissionError(f"Path outside allowed roots: {raw_path}")
31
+
32
+
33
+ def _resolve_repo_root(task: TaskMessage, repositories: dict[str, str] | list[str] | tuple[str, ...]) -> str:
34
+ allowed_path = _resolve_allowed_path(task, repositories)
35
+ return git_utils.resolve_git_repository_path(allowed_path)
36
+
37
+
38
+ def execute_task(task: TaskMessage, repositories: dict[str, str] | list[str] | tuple[str, ...]) -> dict:
39
+ handler = _HANDLERS.get((task.capability, task.action))
40
+ if handler is None:
41
+ raise KeyError(f"Unsupported task: {task.capability}/{task.action}")
42
+ git_utils.clear_cache()
43
+ return handler(task, repositories)
44
+
45
+
46
+ def _fs_tree(task: TaskMessage, repositories) -> dict:
47
+ path = _resolve_allowed_path(task, repositories)
48
+ include_nested = bool(task.params.get("all", False))
49
+ if git_utils.is_git_repository_path(path):
50
+ repo_path = git_utils.resolve_git_repository_path(path)
51
+ return git_utils.get_file_tree(repo_path, changed_only=not include_nested)
52
+ return git_utils.get_workspace_tree(path, include_nested=include_nested)
53
+
54
+
55
+ def _fs_browse(task: TaskMessage, repositories) -> dict:
56
+ path = _resolve_allowed_path(task, repositories)
57
+ return git_utils.browse_workspace(path)
58
+
59
+
60
+ def _fs_read(task: TaskMessage, repositories) -> dict:
61
+ path = _resolve_allowed_path(task, repositories)
62
+ return git_utils.read_workspace_file(path)
63
+
64
+
65
+ def _fs_read_raw(task: TaskMessage, repositories) -> dict:
66
+ path = _resolve_allowed_path(task, repositories)
67
+ max_size = int(task.params.get("max_size_bytes") or task.params.get("max_size") or 0)
68
+ kwargs = {"max_size_bytes": max_size} if max_size > 0 else {}
69
+ return git_utils.read_workspace_file_bytes(path, **kwargs)
70
+
71
+
72
+ def _fs_write(task: TaskMessage, repositories) -> dict:
73
+ path = _resolve_allowed_path(task, repositories)
74
+ return git_utils.write_workspace_file(path, task.params["content"], task.params.get("encoding", "utf-8"))
75
+
76
+
77
+ def _fs_mkdir(task: TaskMessage, repositories) -> dict:
78
+ path = _resolve_allowed_path(task, repositories)
79
+ return git_utils.create_workspace_folder(path, bool(task.params.get("exist_ok", True)))
80
+
81
+
82
+ def _fs_delete(task: TaskMessage, repositories) -> dict:
83
+ path = _resolve_allowed_path(task, repositories)
84
+ for root in _allowed_roots(repositories):
85
+ if path == root:
86
+ raise PermissionError(f"Refusing to delete workspace root: {path}")
87
+ return git_utils.delete_workspace_path(path, bool(task.params.get("recursive", False)))
88
+
89
+
90
+ def _fs_rename(task: TaskMessage, repositories) -> dict:
91
+ src = _resolve_allowed_path(task, repositories)
92
+ raw_destination = task.params.get("to") or task.params.get("new_path")
93
+ if not raw_destination:
94
+ raise ValueError("Missing destination path (params.to)")
95
+ dst = _ensure_allowed_path(str(raw_destination), repositories)
96
+ for root in _allowed_roots(repositories):
97
+ if src == root:
98
+ raise PermissionError(f"Refusing to rename workspace root: {src}")
99
+ return git_utils.rename_workspace_path(src, dst)
100
+
101
+
102
+ def _tree(task: TaskMessage, repositories) -> dict:
103
+ show_all = bool(task.params.get("all", False))
104
+ return git_utils.get_file_tree(_resolve_repo_root(task, repositories), changed_only=not show_all)
105
+
106
+
107
+ def _browse(task: TaskMessage, repositories) -> dict:
108
+ return git_utils.browse_directory(_resolve_repo_root(task, repositories), task.params.get("path", "/"))
109
+
110
+
111
+ def _diff(task: TaskMessage, repositories) -> dict:
112
+ return git_utils.get_file_diff(
113
+ _resolve_repo_root(task, repositories),
114
+ task.params["file"],
115
+ bool(task.params.get("staged", False)),
116
+ int(task.params.get("context", 3)),
117
+ )
118
+
119
+
120
+ def _status(task: TaskMessage, repositories) -> dict:
121
+ return git_utils.get_repo_status(_resolve_repo_root(task, repositories))
122
+
123
+
124
+ def _commits(task: TaskMessage, repositories) -> dict:
125
+ limit = int(task.params.get("limit", 50))
126
+ offset = int(task.params.get("offset", 0))
127
+ repo_path = _resolve_repo_root(task, repositories)
128
+ return {
129
+ "commits": git_utils.get_commit_history(repo_path, limit=limit, offset=offset),
130
+ "limit": limit,
131
+ "offset": offset,
132
+ }
133
+
134
+
135
+ def _commit_details(task: TaskMessage, repositories) -> dict:
136
+ return git_utils.get_commit_details(_resolve_repo_root(task, repositories), task.params["commit_hash"])
137
+
138
+
139
+ def _read_file(task: TaskMessage, repositories) -> dict:
140
+ return git_utils.read_file_content(_resolve_repo_root(task, repositories), task.params["file_path"])
141
+
142
+
143
+ def _write_file(task: TaskMessage, repositories) -> dict:
144
+ return git_utils.write_file_content(
145
+ _resolve_repo_root(task, repositories),
146
+ task.params["file_path"],
147
+ task.params["content"],
148
+ task.params.get("encoding", "utf-8"),
149
+ )
150
+
151
+
152
+ def _stage(task: TaskMessage, repositories) -> dict:
153
+ repo_path = _resolve_repo_root(task, repositories)
154
+ if task.params.get("stage_all"):
155
+ return git_utils.stage_all_files(repo_path)
156
+ return git_utils.stage_file(repo_path, task.params["file_path"])
157
+
158
+
159
+ def _unstage(task: TaskMessage, repositories) -> dict:
160
+ return git_utils.unstage_file(_resolve_repo_root(task, repositories), task.params["file_path"])
161
+
162
+
163
+ def _list_branches(task: TaskMessage, repositories) -> dict:
164
+ return git_utils.list_branches(_resolve_repo_root(task, repositories))
165
+
166
+
167
+ def _create_branch(task: TaskMessage, repositories) -> dict:
168
+ return git_utils.create_branch(_resolve_repo_root(task, repositories), task.params["name"], task.params.get("start_point"))
169
+
170
+
171
+ def _checkout_branch(task: TaskMessage, repositories) -> dict:
172
+ return git_utils.switch_branch(_resolve_repo_root(task, repositories), task.params["name"])
173
+
174
+
175
+ def _delete_branch(task: TaskMessage, repositories) -> dict:
176
+ return git_utils.delete_branch(_resolve_repo_root(task, repositories), task.params["name"], bool(task.params.get("force", False)))
177
+
178
+
179
+ def _merge_branch(task: TaskMessage, repositories) -> dict:
180
+ return git_utils.merge_branch(_resolve_repo_root(task, repositories), task.params["source"], task.params.get("message"))
181
+
182
+
183
+ def _create_commit(task: TaskMessage, repositories) -> dict:
184
+ return git_utils.create_commit(
185
+ _resolve_repo_root(task, repositories),
186
+ task.params["message"],
187
+ task.params.get("author_name"),
188
+ task.params.get("author_email"),
189
+ )
190
+
191
+
192
+ def _list_remotes(task: TaskMessage, repositories) -> dict:
193
+ return git_utils.list_remotes(_resolve_repo_root(task, repositories))
194
+
195
+
196
+ def _add_remote(task: TaskMessage, repositories) -> dict:
197
+ return git_utils.add_remote(
198
+ _resolve_repo_root(task, repositories),
199
+ task.params["name"],
200
+ task.params["url"],
201
+ )
202
+
203
+
204
+ def _remove_remote(task: TaskMessage, repositories) -> dict:
205
+ return git_utils.remove_remote(
206
+ _resolve_repo_root(task, repositories),
207
+ task.params["name"],
208
+ )
209
+
210
+
211
+ def _push_remote(task: TaskMessage, repositories) -> dict:
212
+ return git_utils.push_to_remote(
213
+ _resolve_repo_root(task, repositories),
214
+ task.params["remote"],
215
+ task.params.get("local_branch"),
216
+ task.params.get("remote_branch"),
217
+ bool(task.params.get("set_upstream", False)),
218
+ bool(task.params.get("force", False)),
219
+ )
220
+
221
+
222
+ def _pull_remote(task: TaskMessage, repositories) -> dict:
223
+ return git_utils.pull_from_remote(
224
+ _resolve_repo_root(task, repositories),
225
+ task.params["remote"],
226
+ task.params.get("local_branch"),
227
+ task.params.get("remote_branch"),
228
+ bool(task.params.get("rebase", False)),
229
+ )
230
+
231
+
232
+ _HANDLERS: dict[tuple[str, str], Callable[[TaskMessage, dict[str, str] | list[str] | tuple[str, ...]], dict]] = {
233
+ ("fs.tree", "get"): _fs_tree,
234
+ ("fs.browse", "get"): _fs_browse,
235
+ ("fs.file.read", "get"): _fs_read,
236
+ ("fs.file.read_raw", "get"): _fs_read_raw,
237
+ ("fs.file.write", "put"): _fs_write,
238
+ ("fs.folder.create", "post"): _fs_mkdir,
239
+ ("fs.path.delete", "delete"): _fs_delete,
240
+ ("fs.path.rename", "put"): _fs_rename,
241
+ ("repo.tree", "get"): _tree,
242
+ ("repo.browse", "get"): _browse,
243
+ ("repo.diff", "get"): _diff,
244
+ ("repo.status", "get"): _status,
245
+ ("repo.commits", "list"): _commits,
246
+ ("repo.commit_details", "get"): _commit_details,
247
+ ("repo.file.read", "get"): _read_file,
248
+ ("repo.file.write", "put"): _write_file,
249
+ ("repo.stage", "post"): _stage,
250
+ ("repo.unstage", "post"): _unstage,
251
+ ("repo.branches.list", "get"): _list_branches,
252
+ ("repo.branches.create", "post"): _create_branch,
253
+ ("repo.branches.checkout", "put"): _checkout_branch,
254
+ ("repo.branches.delete", "delete"): _delete_branch,
255
+ ("repo.branches.merge", "post"): _merge_branch,
256
+ ("repo.commit.create", "post"): _create_commit,
257
+ ("repo.remotes.list", "get"): _list_remotes,
258
+ ("repo.remotes.add", "post"): _add_remote,
259
+ ("repo.remotes.remove", "delete"): _remove_remote,
260
+ ("repo.remotes.push", "post"): _push_remote,
261
+ ("repo.remotes.pull", "post"): _pull_remote,
262
+ ("git.remotes.list", "get"): _list_remotes,
263
+ ("git.remotes.add", "post"): _add_remote,
264
+ ("git.remotes.remove", "delete"): _remove_remote,
265
+ ("git.remotes.push", "post"): _push_remote,
266
+ ("git.remotes.pull", "post"): _pull_remote,
267
+ ("git.status", "get"): _status,
268
+ ("git.diff", "get"): _diff,
269
+ ("git.commits", "list"): _commits,
270
+ ("git.commit_details", "get"): _commit_details,
271
+ ("git.stage", "post"): _stage,
272
+ ("git.unstage", "post"): _unstage,
273
+ ("git.branches.list", "get"): _list_branches,
274
+ ("git.branches.create", "post"): _create_branch,
275
+ ("git.branches.checkout", "put"): _checkout_branch,
276
+ ("git.branches.delete", "delete"): _delete_branch,
277
+ ("git.branches.merge", "post"): _merge_branch,
278
+ ("git.commit.create", "post"): _create_commit,
279
+ }
@@ -0,0 +1,179 @@
1
+ from __future__ import annotations
2
+
3
+ import json
4
+ import os
5
+ import platform
6
+ import tomllib
7
+ from dataclasses import dataclass, field
8
+ from importlib import metadata
9
+ from pathlib import Path
10
+
11
+
12
+ DEFAULT_CONFIG_DIRNAME = ".bamboo-coding"
13
+ DEFAULT_CONFIG_FILENAME = "bamboo-coding.toml"
14
+ DEFAULT_CONTROLLER_URL = "ws://127.0.0.1:8100/ws/agents"
15
+ DEFAULT_CLIENT_ID = "bamboo-coding-client"
16
+
17
+
18
+ @dataclass(frozen=True)
19
+ class RepositoryConfig:
20
+ repo_id: str
21
+ name: str
22
+ path: str
23
+
24
+
25
+ @dataclass(frozen=True)
26
+ class AgentConfig:
27
+ controller_url: str
28
+ client_id: str
29
+ hostname: str
30
+ version: str
31
+ token: str
32
+ token_path: Path
33
+ journal_path: Path
34
+ config_path: Path = field(default_factory=Path)
35
+ max_concurrent_tasks: int = 1
36
+ allowed_roots: tuple[str, ...] = field(default_factory=tuple)
37
+ repositories: tuple[RepositoryConfig, ...] = field(default_factory=tuple)
38
+
39
+
40
+ def _first_env(*names: str) -> str | None:
41
+ for name in names:
42
+ value = os.getenv(name)
43
+ if value:
44
+ return value
45
+ return None
46
+
47
+
48
+ def _default_version() -> str:
49
+ try:
50
+ return metadata.version("bamboo-coding")
51
+ except metadata.PackageNotFoundError:
52
+ return "0.1.0"
53
+
54
+
55
+ def _normalize_root(path_value: str) -> str:
56
+ return str(Path(path_value).expanduser().resolve())
57
+
58
+
59
+ def _resolve_path(value: str | None, *, default: Path, base_dir: Path) -> Path:
60
+ if not value:
61
+ return default
62
+ path = Path(value).expanduser()
63
+ if not path.is_absolute():
64
+ path = base_dir / path
65
+ return path.resolve()
66
+
67
+
68
+ def _load_toml(path: Path) -> dict:
69
+ if not path.exists():
70
+ return {}
71
+ with path.open("rb") as handle:
72
+ data = tomllib.load(handle)
73
+ return data if isinstance(data, dict) else {}
74
+
75
+
76
+ def _coerce_allowed_roots(value: object) -> tuple[str, ...]:
77
+ if isinstance(value, str):
78
+ items = [item for item in value.split(os.pathsep) if item.strip()]
79
+ elif isinstance(value, (list, tuple)):
80
+ items = [str(item) for item in value if str(item).strip()]
81
+ else:
82
+ items = []
83
+ return tuple(_normalize_root(item) for item in items)
84
+
85
+
86
+ def _load_stored_token(token_path: Path) -> str:
87
+ if not token_path.exists():
88
+ return ""
89
+ try:
90
+ payload = json.loads(token_path.read_text(encoding="utf-8"))
91
+ except (OSError, json.JSONDecodeError):
92
+ return ""
93
+ token = payload.get("token")
94
+ return token if isinstance(token, str) else ""
95
+
96
+
97
+ def _persist_token(token_path: Path, token: str) -> None:
98
+ token_path.parent.mkdir(parents=True, exist_ok=True)
99
+ token_path.write_text(json.dumps({"token": token}), encoding="utf-8")
100
+
101
+
102
+ def _load_legacy_token(default_token_path: Path) -> str:
103
+ legacy_path = Path.cwd() / ".agent" / "token.json"
104
+ token = _load_stored_token(legacy_path)
105
+ if token and not default_token_path.exists():
106
+ _persist_token(default_token_path, token)
107
+ return token
108
+
109
+
110
+ def load_config() -> AgentConfig:
111
+ config_home = Path(_first_env("BAMBOO_CODING_HOME") or (Path.home() / DEFAULT_CONFIG_DIRNAME)).expanduser().resolve()
112
+ config_path = _resolve_path(
113
+ _first_env("BAMBOO_CODING_CONFIG_PATH"),
114
+ default=config_home / DEFAULT_CONFIG_FILENAME,
115
+ base_dir=config_home,
116
+ )
117
+ raw_config = _load_toml(config_path)
118
+ client_config = raw_config.get("client") if isinstance(raw_config.get("client"), dict) else {}
119
+ storage_config = raw_config.get("storage") if isinstance(raw_config.get("storage"), dict) else {}
120
+
121
+ storage_home_value = _first_env("BAMBOO_CODING_HOME", "AGENT_HOME")
122
+ storage_home = Path(storage_home_value).expanduser().resolve() if storage_home_value else config_home
123
+
124
+ controller_url = _first_env("BAMBOO_CODING_CONTROLLER_URL", "CONTROLLER_URL") or str(
125
+ client_config.get("controller_url") or DEFAULT_CONTROLLER_URL
126
+ )
127
+ client_id = _first_env("BAMBOO_CODING_CLIENT_ID", "AGENT_CLIENT_ID") or str(
128
+ client_config.get("client_id") or platform.node() or DEFAULT_CLIENT_ID
129
+ )
130
+ hostname = _first_env("BAMBOO_CODING_HOSTNAME", "AGENT_HOSTNAME") or str(
131
+ client_config.get("hostname") or platform.node() or client_id
132
+ )
133
+ version = _first_env("BAMBOO_CODING_VERSION", "AGENT_VERSION") or str(client_config.get("version") or _default_version())
134
+ allowed_roots = _coerce_allowed_roots(
135
+ _first_env("BAMBOO_CODING_ALLOWED_ROOTS", "AGENT_ALLOWED_ROOTS", "AGENT_ROOT", "AGENT_REPO_PATH")
136
+ or client_config.get("allowed_roots")
137
+ )
138
+ max_concurrent_tasks = int(
139
+ _first_env("BAMBOO_CODING_MAX_CONCURRENT_TASKS", "AGENT_MAX_CONCURRENT_TASKS")
140
+ or client_config.get("max_concurrent_tasks")
141
+ or 1
142
+ )
143
+
144
+ token_path = _resolve_path(
145
+ _first_env("BAMBOO_CODING_TOKEN_PATH", "AGENT_TOKEN_PATH") or storage_config.get("token_path"),
146
+ default=storage_home / "token.json",
147
+ base_dir=config_path.parent,
148
+ )
149
+ journal_path = _resolve_path(
150
+ _first_env("BAMBOO_CODING_JOURNAL_PATH", "AGENT_JOURNAL_PATH") or storage_config.get("journal_path"),
151
+ default=storage_home / "journal.db",
152
+ base_dir=config_path.parent,
153
+ )
154
+
155
+ token = _first_env("BAMBOO_CODING_TOKEN", "AGENT_TOKEN") or _load_stored_token(token_path)
156
+ if not token and not _first_env("BAMBOO_CODING_TOKEN_PATH", "AGENT_TOKEN_PATH", "AGENT_HOME"):
157
+ token = _load_legacy_token(token_path)
158
+
159
+ repositories = tuple(
160
+ RepositoryConfig(
161
+ repo_id=root,
162
+ name=Path(root).name or root,
163
+ path=root,
164
+ )
165
+ for root in allowed_roots
166
+ )
167
+ return AgentConfig(
168
+ controller_url=controller_url,
169
+ client_id=client_id,
170
+ hostname=hostname,
171
+ version=version,
172
+ token=token,
173
+ token_path=token_path,
174
+ journal_path=journal_path,
175
+ config_path=config_path,
176
+ max_concurrent_tasks=max_concurrent_tasks,
177
+ allowed_roots=allowed_roots,
178
+ repositories=repositories,
179
+ )
@@ -0,0 +1,112 @@
1
+ from __future__ import annotations
2
+
3
+ import asyncio
4
+ import json
5
+ from typing import Awaitable, Callable
6
+
7
+ from websockets import connect
8
+ from websockets.exceptions import ConnectionClosed
9
+
10
+ from bamboo_coding.shared.protocol import (
11
+ AgentRegisterMessage,
12
+ ClientStatusMessage,
13
+ TaskEventMessage,
14
+ TaskMessage,
15
+ TaskQueryResultMessage,
16
+ TaskResultMessage,
17
+ TerminalCloseMessage,
18
+ TerminalInputMessage,
19
+ TerminalHeartbeatMessage,
20
+ TerminalOpenMessage,
21
+ TerminalOutputMessage,
22
+ TerminalResizeMessage,
23
+ TerminalStatusMessage,
24
+ TerminalClosedMessage,
25
+ )
26
+
27
+
28
+ TerminalControlMessage = TerminalOpenMessage | TerminalInputMessage | TerminalResizeMessage | TerminalCloseMessage
29
+
30
+
31
+ class ControllerConnection:
32
+ def __init__(self, controller_url: str, token: str = ''):
33
+ self.controller_url = controller_url
34
+ self.token = token
35
+ self.websocket = None
36
+ self._send_lock = asyncio.Lock()
37
+
38
+ async def connect(self):
39
+ headers = {}
40
+ if self.token:
41
+ headers['Authorization'] = f'Bearer {self.token}'
42
+ self.websocket = await connect(self.controller_url, extra_headers=headers or None, ping_interval=20, ping_timeout=20)
43
+ return self.websocket
44
+
45
+ async def send_register(self, message: AgentRegisterMessage) -> dict:
46
+ await self._send(message.model_dump())
47
+ return await self._receive()
48
+
49
+ async def send_status(self, message: ClientStatusMessage) -> None:
50
+ await self._send(message.model_dump())
51
+
52
+ async def send_event(self, message: TaskEventMessage) -> None:
53
+ await self._send(message.model_dump())
54
+
55
+ async def send_result(self, message: TaskResultMessage) -> None:
56
+ await self._send(message.model_dump())
57
+
58
+ async def send_query_result(self, message: TaskQueryResultMessage) -> None:
59
+ await self._send(message.model_dump())
60
+
61
+ async def send_terminal_status(self, message: TerminalStatusMessage) -> None:
62
+ await self._send(message.model_dump())
63
+
64
+ async def send_terminal_output(self, message: TerminalOutputMessage) -> None:
65
+ await self._send(message.model_dump())
66
+
67
+ async def send_terminal_heartbeat(self, message: TerminalHeartbeatMessage) -> None:
68
+ await self._send(message.model_dump())
69
+
70
+ async def send_terminal_closed(self, message: TerminalClosedMessage) -> None:
71
+ await self._send(message.model_dump())
72
+
73
+ async def listen(
74
+ self,
75
+ on_task: Callable[[TaskMessage], Awaitable[None]],
76
+ on_terminal_message: Callable[[TerminalControlMessage], Awaitable[None]] | None = None,
77
+ ) -> None:
78
+ if self.websocket is None:
79
+ raise RuntimeError('connection not established')
80
+ try:
81
+ async for raw in self.websocket:
82
+ data = json.loads(raw)
83
+ message_type = data.get('type')
84
+ if message_type == 'task':
85
+ await on_task(TaskMessage(**data))
86
+ elif on_terminal_message is not None and message_type == 'terminal_open':
87
+ await on_terminal_message(TerminalOpenMessage(**data))
88
+ elif on_terminal_message is not None and message_type == 'terminal_input':
89
+ await on_terminal_message(TerminalInputMessage(**data))
90
+ elif on_terminal_message is not None and message_type == 'terminal_resize':
91
+ await on_terminal_message(TerminalResizeMessage(**data))
92
+ elif on_terminal_message is not None and message_type == 'terminal_close':
93
+ await on_terminal_message(TerminalCloseMessage(**data))
94
+ except ConnectionClosed:
95
+ return
96
+
97
+ async def close(self) -> None:
98
+ if self.websocket is not None:
99
+ await self.websocket.close()
100
+ self.websocket = None
101
+
102
+ async def _send(self, payload: dict) -> None:
103
+ if self.websocket is None:
104
+ raise RuntimeError('connection not established')
105
+ async with self._send_lock:
106
+ await self.websocket.send(json.dumps(payload))
107
+
108
+ async def _receive(self) -> dict:
109
+ if self.websocket is None:
110
+ raise RuntimeError('connection not established')
111
+ raw = await self.websocket.recv()
112
+ return json.loads(raw)
@@ -0,0 +1 @@
1
+ from agent.bamboo_git_agent.journal import * # noqa: F401,F403