splitagent 0.0.3__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.
- splitagent/__init__.py +8 -0
- splitagent/__main__.py +6 -0
- splitagent/agents/__init__.py +10 -0
- splitagent/agents/base.py +477 -0
- splitagent/agents/blue.py +57 -0
- splitagent/agents/chat.py +60 -0
- splitagent/agents/prompts.py +462 -0
- splitagent/agents/red.py +75 -0
- splitagent/cli.py +701 -0
- splitagent/config.py +697 -0
- splitagent/core/__init__.py +19 -0
- splitagent/core/bus.py +62 -0
- splitagent/core/context.py +587 -0
- splitagent/core/context_manager.py +381 -0
- splitagent/core/engine.py +424 -0
- splitagent/core/models.py +310 -0
- splitagent/core/proc.py +73 -0
- splitagent/core/sandbox.py +184 -0
- splitagent/core/toolbox.py +520 -0
- splitagent/core/workspace.py +420 -0
- splitagent/desktop/__init__.py +7 -0
- splitagent/desktop/api.py +525 -0
- splitagent/desktop/app.py +1131 -0
- splitagent/desktop/web/app.js +3067 -0
- splitagent/desktop/web/assets/Inter.ttf +0 -0
- splitagent/desktop/web/assets/JetBrainsMonoNerdFontMono-Regular.woff2 +0 -0
- splitagent/desktop/web/index.html +760 -0
- splitagent/desktop/web/styles.css +1612 -0
- splitagent/errors.py +27 -0
- splitagent/llm/__init__.py +8 -0
- splitagent/llm/client.py +488 -0
- splitagent/llm/types.py +172 -0
- splitagent/report/__init__.py +9 -0
- splitagent/report/cvss.py +93 -0
- splitagent/report/generator.py +733 -0
- splitagent/tools/__init__.py +8 -0
- splitagent/tools/base.py +135 -0
- splitagent/tools/defense.py +475 -0
- splitagent/tools/exploit.py +318 -0
- splitagent/tools/http_pool.py +109 -0
- splitagent/tools/knowledge.py +376 -0
- splitagent/tools/recon.py +182 -0
- splitagent/tools/registry.py +62 -0
- splitagent/tools/validate.py +908 -0
- splitagent/tools/web.py +386 -0
- splitagent/tools/workspace_tools.py +411 -0
- splitagent/ui/__init__.py +5 -0
- splitagent/ui/app.py +389 -0
- splitagent/ui/stream.py +234 -0
- splitagent/ui/theme.py +72 -0
- splitagent-0.0.3.dist-info/METADATA +987 -0
- splitagent-0.0.3.dist-info/RECORD +56 -0
- splitagent-0.0.3.dist-info/WHEEL +5 -0
- splitagent-0.0.3.dist-info/entry_points.txt +2 -0
- splitagent-0.0.3.dist-info/licenses/LICENSE +21 -0
- splitagent-0.0.3.dist-info/top_level.txt +1 -0
|
@@ -0,0 +1,1131 @@
|
|
|
1
|
+
"""Native desktop shell: a frameless pywebview window backed by the engine."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import asyncio
|
|
6
|
+
import copy
|
|
7
|
+
import json
|
|
8
|
+
import re
|
|
9
|
+
import threading
|
|
10
|
+
import time
|
|
11
|
+
from datetime import datetime, timezone
|
|
12
|
+
from pathlib import Path
|
|
13
|
+
from typing import Any
|
|
14
|
+
|
|
15
|
+
from splitagent.agents.chat import ChatAgent
|
|
16
|
+
from splitagent.config import (
|
|
17
|
+
PROVIDER_PRESETS,
|
|
18
|
+
GlobalConfig,
|
|
19
|
+
LLMSettings,
|
|
20
|
+
ProjectConfig,
|
|
21
|
+
ProviderConfig,
|
|
22
|
+
config_home,
|
|
23
|
+
ensure_provider,
|
|
24
|
+
project_config_path,
|
|
25
|
+
save_global_config,
|
|
26
|
+
save_project_config,
|
|
27
|
+
sessions_dir,
|
|
28
|
+
settings_for_provider,
|
|
29
|
+
)
|
|
30
|
+
from splitagent.core.bus import Event, EventBus
|
|
31
|
+
from splitagent.core.context import SharedContext
|
|
32
|
+
from splitagent.core.engine import Engine
|
|
33
|
+
from splitagent.core.toolbox import (
|
|
34
|
+
Toolbox,
|
|
35
|
+
build_command,
|
|
36
|
+
estimated_size,
|
|
37
|
+
resolve_mode,
|
|
38
|
+
)
|
|
39
|
+
from splitagent.core.workspace import workspace_for
|
|
40
|
+
from splitagent.desktop.api import JsApi
|
|
41
|
+
from splitagent.llm.client import LLMClient
|
|
42
|
+
from splitagent.llm.types import ChatMessage
|
|
43
|
+
|
|
44
|
+
WEB_DIR = Path(__file__).parent / "web"
|
|
45
|
+
FLUSH_INTERVAL = 0.06
|
|
46
|
+
# The planner is one short call; past this we use the deterministic fallback.
|
|
47
|
+
PLANNER_TIMEOUT_SECONDS = 15.0
|
|
48
|
+
|
|
49
|
+
|
|
50
|
+
def _extract_json_object(text: str) -> dict[str, Any]:
|
|
51
|
+
"""Pull the first JSON object out of a model reply.
|
|
52
|
+
|
|
53
|
+
Providers without a JSON mode sometimes wrap the object in prose or a
|
|
54
|
+
fenced code block; we tolerate all of that and fail soft (empty dict).
|
|
55
|
+
"""
|
|
56
|
+
if not text:
|
|
57
|
+
return {}
|
|
58
|
+
fenced = re.search(r"```(?:json)?\s*(\{.*?\})\s*```", text, re.DOTALL)
|
|
59
|
+
candidate = fenced.group(1) if fenced else None
|
|
60
|
+
if candidate is None:
|
|
61
|
+
start = text.find("{")
|
|
62
|
+
end = text.rfind("}")
|
|
63
|
+
if start == -1 or end <= start:
|
|
64
|
+
return {}
|
|
65
|
+
candidate = text[start : end + 1]
|
|
66
|
+
try:
|
|
67
|
+
parsed = json.loads(candidate)
|
|
68
|
+
except (ValueError, TypeError):
|
|
69
|
+
return {}
|
|
70
|
+
return parsed if isinstance(parsed, dict) else {}
|
|
71
|
+
|
|
72
|
+
|
|
73
|
+
def _compose_brief_instructions(brief: dict[str, Any]) -> str:
|
|
74
|
+
"""Turn the audit brief into operator instructions for the agents.
|
|
75
|
+
|
|
76
|
+
This is what makes the brief actually reach the Red/Blue agents instead of
|
|
77
|
+
being lost as a project name.
|
|
78
|
+
"""
|
|
79
|
+
if not isinstance(brief, dict):
|
|
80
|
+
return ""
|
|
81
|
+
lines: list[str] = []
|
|
82
|
+
label = {
|
|
83
|
+
"objective": "Objective",
|
|
84
|
+
"crown_jewels": "What matters most",
|
|
85
|
+
"authorization": "Authorization",
|
|
86
|
+
"goal": "Engagement goal",
|
|
87
|
+
"exclude_areas": "Areas/endpoints to exclude",
|
|
88
|
+
"window": "Window / rate constraints",
|
|
89
|
+
"noise": "Noise profile",
|
|
90
|
+
"notes": "Additional notes",
|
|
91
|
+
}
|
|
92
|
+
for key, title in label.items():
|
|
93
|
+
value = brief.get(key)
|
|
94
|
+
if value is None:
|
|
95
|
+
continue
|
|
96
|
+
if isinstance(value, list):
|
|
97
|
+
value = ", ".join(str(v) for v in value if str(v).strip())
|
|
98
|
+
value = str(value).strip()
|
|
99
|
+
if value:
|
|
100
|
+
lines.append(f"- {title}: {value}")
|
|
101
|
+
scope_doc = str(brief.get("scope_document") or "").strip()
|
|
102
|
+
if scope_doc:
|
|
103
|
+
lines.append("- Operator scope document:")
|
|
104
|
+
lines.append(scope_doc)
|
|
105
|
+
if not lines:
|
|
106
|
+
return ""
|
|
107
|
+
return "=== ENGAGEMENT BRIEF ===\n" + "\n".join(lines)
|
|
108
|
+
|
|
109
|
+
|
|
110
|
+
async def _fetch_models_async(settings: LLMSettings) -> list[dict[str, str]]:
|
|
111
|
+
async with LLMClient(settings) as client:
|
|
112
|
+
return await client.list_models()
|
|
113
|
+
|
|
114
|
+
|
|
115
|
+
def _event_dict(event: Event) -> dict[str, Any]:
|
|
116
|
+
return {
|
|
117
|
+
"type": event.type,
|
|
118
|
+
"agent": event.agent,
|
|
119
|
+
"data": event.data,
|
|
120
|
+
"ts": event.ts,
|
|
121
|
+
}
|
|
122
|
+
|
|
123
|
+
|
|
124
|
+
class DesktopApp:
|
|
125
|
+
"""Owns the webview window, the JS bridge and the audit worker."""
|
|
126
|
+
|
|
127
|
+
def __init__(self, global_config: GlobalConfig, project: ProjectConfig) -> None:
|
|
128
|
+
self.global_config = global_config
|
|
129
|
+
self.project = project
|
|
130
|
+
self.window: Any = None
|
|
131
|
+
self.api = JsApi(self)
|
|
132
|
+
self.session: SharedContext | None = None
|
|
133
|
+
self.report_paths: list[str] = []
|
|
134
|
+
self._audit_thread: threading.Thread | None = None
|
|
135
|
+
self._loop: asyncio.AbstractEventLoop | None = None
|
|
136
|
+
self._task: asyncio.Task[Any] | None = None
|
|
137
|
+
self._queue: list[dict[str, Any]] = []
|
|
138
|
+
self._last_flush = 0.0
|
|
139
|
+
self._stop = threading.Event()
|
|
140
|
+
self._maximized = False
|
|
141
|
+
# copilot chat state
|
|
142
|
+
self._chat_context: SharedContext | None = None
|
|
143
|
+
self.workspace: Any = None
|
|
144
|
+
self._chat_history: list[ChatMessage] = []
|
|
145
|
+
self._chat_thread: threading.Thread | None = None
|
|
146
|
+
self._chat_loop: asyncio.AbstractEventLoop | None = None
|
|
147
|
+
self._chat_task: asyncio.Task[Any] | None = None
|
|
148
|
+
self._chat_events: list[dict[str, Any]] = []
|
|
149
|
+
self._chat_done = False
|
|
150
|
+
self._toolbox_thread: threading.Thread | None = None
|
|
151
|
+
self._queue_lock = threading.Lock()
|
|
152
|
+
|
|
153
|
+
# -- lifecycle --------------------------------------------------------- #
|
|
154
|
+
def run(self) -> None:
|
|
155
|
+
import webview
|
|
156
|
+
|
|
157
|
+
index = (WEB_DIR / "index.html").as_uri()
|
|
158
|
+
self.window = webview.create_window(
|
|
159
|
+
"SplitAgent",
|
|
160
|
+
url=index,
|
|
161
|
+
js_api=self.api,
|
|
162
|
+
width=1460,
|
|
163
|
+
height=940,
|
|
164
|
+
min_size=(760, 540),
|
|
165
|
+
frameless=True,
|
|
166
|
+
easy_drag=False,
|
|
167
|
+
background_color="#161616",
|
|
168
|
+
text_select=True,
|
|
169
|
+
confirm_close=False,
|
|
170
|
+
)
|
|
171
|
+
self.window.events.closed += self._on_closed
|
|
172
|
+
storage = config_home() / "webview"
|
|
173
|
+
storage.mkdir(parents=True, exist_ok=True)
|
|
174
|
+
webview.start(
|
|
175
|
+
self._keepalive,
|
|
176
|
+
gui="edgechromium",
|
|
177
|
+
debug=False,
|
|
178
|
+
private_mode=False,
|
|
179
|
+
storage_path=str(storage),
|
|
180
|
+
)
|
|
181
|
+
|
|
182
|
+
def _keepalive(self) -> None:
|
|
183
|
+
# Runs in a background thread for the lifetime of the window.
|
|
184
|
+
while not self._stop.is_set():
|
|
185
|
+
time.sleep(0.25)
|
|
186
|
+
|
|
187
|
+
def _on_closed(self) -> None:
|
|
188
|
+
self._stop.set()
|
|
189
|
+
|
|
190
|
+
def is_running(self) -> bool:
|
|
191
|
+
return self._audit_thread is not None and self._audit_thread.is_alive()
|
|
192
|
+
|
|
193
|
+
# -- JS bridge helpers ------------------------------------------------- #
|
|
194
|
+
# Threads that feed the bridge: the audit worker, the chat worker and the
|
|
195
|
+
# toolbox worker all push into the same queue. Every read/modify is done
|
|
196
|
+
# under a lock so a batch can never be duplicated or dropped mid-flush.
|
|
197
|
+
def _emit(self, batch: list[dict[str, Any]]) -> None:
|
|
198
|
+
if not batch or self.window is None:
|
|
199
|
+
return
|
|
200
|
+
payload = json.dumps(batch, ensure_ascii=False, default=str)
|
|
201
|
+
script = f"window.SplitAgent && window.SplitAgent.emit({payload});"
|
|
202
|
+
try:
|
|
203
|
+
self.window.evaluate_js(script)
|
|
204
|
+
except Exception: # pragma: no cover - window may be closing
|
|
205
|
+
pass
|
|
206
|
+
|
|
207
|
+
def _emit_chat(self, batch: list[dict[str, Any]]) -> None:
|
|
208
|
+
"""Chat events go only to the polled buffer, never the live bridge.
|
|
209
|
+
|
|
210
|
+
Delivering them both ways duplicated every event (the bridge pushed it
|
|
211
|
+
and the poll replayed it). The front-end polls ``chat_state`` for chat,
|
|
212
|
+
so this is the single channel.
|
|
213
|
+
"""
|
|
214
|
+
if batch:
|
|
215
|
+
with self._queue_lock:
|
|
216
|
+
self._chat_events.extend(batch)
|
|
217
|
+
|
|
218
|
+
def _push(self, event: dict[str, Any]) -> None:
|
|
219
|
+
with self._queue_lock:
|
|
220
|
+
self._queue.append(event)
|
|
221
|
+
|
|
222
|
+
def _flush(self, force: bool = False) -> None:
|
|
223
|
+
now = time.monotonic()
|
|
224
|
+
with self._queue_lock:
|
|
225
|
+
if not self._queue:
|
|
226
|
+
return
|
|
227
|
+
if not force and (now - self._last_flush) < FLUSH_INTERVAL:
|
|
228
|
+
return
|
|
229
|
+
self._last_flush = now
|
|
230
|
+
batch = self._queue[:]
|
|
231
|
+
del self._queue[:]
|
|
232
|
+
self._emit(batch)
|
|
233
|
+
|
|
234
|
+
# -- audit ------------------------------------------------------------- #
|
|
235
|
+
def start_audit(self, options: dict[str, Any]) -> dict[str, Any]:
|
|
236
|
+
if self.is_running():
|
|
237
|
+
return {"ok": False, "error": "An audit is already running."}
|
|
238
|
+
self.report_paths = []
|
|
239
|
+
self.session = None
|
|
240
|
+
self._audit_thread = threading.Thread(
|
|
241
|
+
target=self._audit_entry, args=(options,), daemon=True
|
|
242
|
+
)
|
|
243
|
+
self._audit_thread.start()
|
|
244
|
+
return {"ok": True}
|
|
245
|
+
|
|
246
|
+
def stop_audit(self) -> dict[str, Any]:
|
|
247
|
+
if self._loop is not None and self._task is not None:
|
|
248
|
+
try:
|
|
249
|
+
self._loop.call_soon_threadsafe(self._task.cancel)
|
|
250
|
+
except RuntimeError:
|
|
251
|
+
pass
|
|
252
|
+
return {"ok": True}
|
|
253
|
+
return {"ok": False, "error": "No audit running."}
|
|
254
|
+
|
|
255
|
+
def _audit_entry(self, options: dict[str, Any]) -> None:
|
|
256
|
+
try:
|
|
257
|
+
asyncio.run(self._audit_async(options))
|
|
258
|
+
except asyncio.CancelledError:
|
|
259
|
+
self._push({"type": "error", "agent": "core", "data": {"text": "Audit cancelled."}})
|
|
260
|
+
self._flush(force=True)
|
|
261
|
+
self._emit([{"type": "audit.end", "data": {"ok": False, "error": "cancelled"}}])
|
|
262
|
+
except Exception as exc:
|
|
263
|
+
self._push(
|
|
264
|
+
{
|
|
265
|
+
"type": "error",
|
|
266
|
+
"agent": "core",
|
|
267
|
+
"data": {"text": f"{type(exc).__name__}: {exc}"},
|
|
268
|
+
}
|
|
269
|
+
)
|
|
270
|
+
self._flush(force=True)
|
|
271
|
+
self._emit(
|
|
272
|
+
[
|
|
273
|
+
{
|
|
274
|
+
"type": "audit.end",
|
|
275
|
+
"data": {"ok": False, "error": str(exc)},
|
|
276
|
+
}
|
|
277
|
+
]
|
|
278
|
+
)
|
|
279
|
+
|
|
280
|
+
async def _audit_async(self, options: dict[str, Any]) -> None:
|
|
281
|
+
project = self._build_project(options)
|
|
282
|
+
self.project = project
|
|
283
|
+
|
|
284
|
+
bus = EventBus()
|
|
285
|
+
bus.subscribe(lambda event: self._push(_event_dict(event)))
|
|
286
|
+
engine = Engine(self.global_config, project, bus=bus)
|
|
287
|
+
|
|
288
|
+
self._loop = asyncio.get_running_loop()
|
|
289
|
+
self._emit(
|
|
290
|
+
[
|
|
291
|
+
{
|
|
292
|
+
"type": "audit.start",
|
|
293
|
+
"data": {
|
|
294
|
+
"target": project.target.url,
|
|
295
|
+
"model": self.global_config.llm.model,
|
|
296
|
+
"provider": self.global_config.llm.provider,
|
|
297
|
+
"rounds": project.run.rounds,
|
|
298
|
+
},
|
|
299
|
+
}
|
|
300
|
+
]
|
|
301
|
+
)
|
|
302
|
+
self._task = asyncio.ensure_future(engine.run())
|
|
303
|
+
while not self._task.done():
|
|
304
|
+
await asyncio.sleep(0.04)
|
|
305
|
+
self._flush()
|
|
306
|
+
context = self._task.result()
|
|
307
|
+
self.session = context
|
|
308
|
+
self._flush(force=True)
|
|
309
|
+
|
|
310
|
+
report_paths: list[str] = []
|
|
311
|
+
try:
|
|
312
|
+
from splitagent.report.generator import write_reports
|
|
313
|
+
|
|
314
|
+
paths = write_reports(context.state, project.report, Path(project.report.output_dir))
|
|
315
|
+
report_paths = [str(p.resolve()) for p in paths]
|
|
316
|
+
except Exception as exc:
|
|
317
|
+
self._push({"type": "error", "agent": "core", "data": {"text": f"report: {exc}"}})
|
|
318
|
+
self.report_paths = report_paths
|
|
319
|
+
# Counting completed audits feeds the dashboard header.
|
|
320
|
+
self.global_config.ui.audits_completed += 1
|
|
321
|
+
save_global_config(self.global_config)
|
|
322
|
+
self._emit(
|
|
323
|
+
[
|
|
324
|
+
{
|
|
325
|
+
"type": "audit.end",
|
|
326
|
+
"data": {
|
|
327
|
+
"ok": True,
|
|
328
|
+
"session": context.state.id,
|
|
329
|
+
"reports": report_paths,
|
|
330
|
+
"summary": context.summary_dict(),
|
|
331
|
+
},
|
|
332
|
+
}
|
|
333
|
+
]
|
|
334
|
+
)
|
|
335
|
+
|
|
336
|
+
# -- engagement planner (audit brief) ---------------------------------- #
|
|
337
|
+
def build_audit_plan(self, brief: dict[str, Any] | None = None) -> dict[str, Any]:
|
|
338
|
+
"""Turn an operator brief into a plan, with the model doing the adapting.
|
|
339
|
+
|
|
340
|
+
The AI writes the plan (phases, techniques, cautions) because a fixed
|
|
341
|
+
template cannot fit every environment. Hard facts — scope and
|
|
342
|
+
out_of_scope — come from the brief and are enforced afterwards: the
|
|
343
|
+
plan is sanitised so a hallucinated host can never enter scope. If the
|
|
344
|
+
model fails or is slow, a deterministic fallback plan is returned so the
|
|
345
|
+
audit always has something to approve.
|
|
346
|
+
"""
|
|
347
|
+
brief = brief or {}
|
|
348
|
+
fallback = self._fallback_plan(brief)
|
|
349
|
+
try:
|
|
350
|
+
plan = asyncio.run(self._plan_async(brief))
|
|
351
|
+
except Exception as exc:
|
|
352
|
+
fallback["source"] = "fallback"
|
|
353
|
+
fallback["note"] = f"planner unavailable: {type(exc).__name__}"
|
|
354
|
+
return fallback
|
|
355
|
+
if not plan:
|
|
356
|
+
fallback["source"] = "fallback"
|
|
357
|
+
fallback["note"] = "planner returned no usable plan"
|
|
358
|
+
return fallback
|
|
359
|
+
return self._sanitise_plan(plan, brief)
|
|
360
|
+
|
|
361
|
+
async def _plan_async(self, brief: dict[str, Any]) -> dict[str, Any]:
|
|
362
|
+
from splitagent.agents.prompts import build_audit_brief_prompt
|
|
363
|
+
|
|
364
|
+
project = self.project
|
|
365
|
+
context = SharedContext.create(
|
|
366
|
+
target=project.target.url
|
|
367
|
+
or ", ".join(project.target.effective_hosts())
|
|
368
|
+
or "unspecified",
|
|
369
|
+
target_kind=project.target.kind,
|
|
370
|
+
scope=project.target.scope or project.target.effective_hosts(),
|
|
371
|
+
name="planner",
|
|
372
|
+
model=self.global_config.llm.model,
|
|
373
|
+
provider=self.global_config.llm.provider,
|
|
374
|
+
)
|
|
375
|
+
system = build_audit_brief_prompt(project, context)
|
|
376
|
+
user = "=== BRIEF ===\n" + json.dumps(brief, ensure_ascii=False, indent=2)
|
|
377
|
+
async with LLMClient(self.global_config.llm) as client:
|
|
378
|
+
message, _usage = await asyncio.wait_for(
|
|
379
|
+
client.complete(
|
|
380
|
+
[
|
|
381
|
+
ChatMessage(role="system", content=system),
|
|
382
|
+
ChatMessage(role="user", content=user),
|
|
383
|
+
]
|
|
384
|
+
),
|
|
385
|
+
timeout=PLANNER_TIMEOUT_SECONDS,
|
|
386
|
+
)
|
|
387
|
+
return _extract_json_object(message.content)
|
|
388
|
+
|
|
389
|
+
def _sanitise_plan(self, plan: dict[str, Any], brief: dict[str, Any]) -> dict[str, Any]:
|
|
390
|
+
"""Keep the AI's judgement but enforce the operator's hard facts."""
|
|
391
|
+
|
|
392
|
+
def _as_list(value: Any) -> list[str]:
|
|
393
|
+
if isinstance(value, list):
|
|
394
|
+
return [str(v).strip() for v in value if str(v).strip()]
|
|
395
|
+
if isinstance(value, str):
|
|
396
|
+
return [p.strip() for p in value.split(",") if p.strip()]
|
|
397
|
+
return []
|
|
398
|
+
|
|
399
|
+
out_of_scope = {h.lower() for h in _as_list(brief.get("out_of_scope"))}
|
|
400
|
+
scope = _as_list(brief.get("scope")) or _as_list(plan.get("scope"))
|
|
401
|
+
# A host can never be both in scope and explicitly excluded.
|
|
402
|
+
scope = [h for h in scope if h.lower() not in out_of_scope]
|
|
403
|
+
noise = str(plan.get("noise") or brief.get("noise") or "normal").lower()
|
|
404
|
+
if noise not in ("stealth", "normal", "aggressive"):
|
|
405
|
+
noise = "normal"
|
|
406
|
+
return {
|
|
407
|
+
"objective": str(plan.get("objective") or brief.get("objective") or "Security audit"),
|
|
408
|
+
"scope": scope,
|
|
409
|
+
"out_of_scope": _as_list(brief.get("out_of_scope"))
|
|
410
|
+
or _as_list(plan.get("out_of_scope")),
|
|
411
|
+
"phases": _as_list(plan.get("phases")),
|
|
412
|
+
"techniques": _as_list(plan.get("techniques")),
|
|
413
|
+
"cautions": _as_list(plan.get("cautions")),
|
|
414
|
+
"noise": noise,
|
|
415
|
+
"source": "ai",
|
|
416
|
+
"awaiting_approval": True,
|
|
417
|
+
}
|
|
418
|
+
|
|
419
|
+
def _fallback_plan(self, brief: dict[str, Any]) -> dict[str, Any]:
|
|
420
|
+
"""A deterministic plan used when the model is unavailable."""
|
|
421
|
+
kind = str(brief.get("kind") or self.project.target.kind or "web")
|
|
422
|
+
goal = str(brief.get("goal") or "full audit")
|
|
423
|
+
common = ["reconnaissance", "enumeration", "targeted validation", "report"]
|
|
424
|
+
by_kind = {
|
|
425
|
+
"web": [
|
|
426
|
+
"map the web surface (crawl, routes, forms)",
|
|
427
|
+
"audit authentication and session handling",
|
|
428
|
+
"test input-handling endpoints (injection, XSS)",
|
|
429
|
+
],
|
|
430
|
+
"api": [
|
|
431
|
+
"enumerate API endpoints and schemas",
|
|
432
|
+
"audit object-level authorisation (IDOR)",
|
|
433
|
+
"test authentication, rate limiting and input validation",
|
|
434
|
+
],
|
|
435
|
+
"network": [
|
|
436
|
+
"service/version discovery",
|
|
437
|
+
"sweep all ports for exposed services",
|
|
438
|
+
"validate known CVEs on discovered versions",
|
|
439
|
+
],
|
|
440
|
+
"repo": [
|
|
441
|
+
"inventory the repository and dependencies",
|
|
442
|
+
"review auth, secrets and injection sinks",
|
|
443
|
+
"check dependency vulnerabilities",
|
|
444
|
+
],
|
|
445
|
+
}
|
|
446
|
+
return {
|
|
447
|
+
"objective": str(brief.get("objective") or f"{goal} of the {kind} target"),
|
|
448
|
+
"scope": [str(s).strip() for s in brief.get("scope", []) if str(s).strip()],
|
|
449
|
+
"out_of_scope": [
|
|
450
|
+
str(s).strip() for s in brief.get("out_of_scope", []) if str(s).strip()
|
|
451
|
+
],
|
|
452
|
+
"phases": by_kind.get(kind, common),
|
|
453
|
+
"techniques": [],
|
|
454
|
+
"cautions": [
|
|
455
|
+
"Stay strictly within the authorised scope.",
|
|
456
|
+
"Non-destructive testing only.",
|
|
457
|
+
]
|
|
458
|
+
+ ([str(brief["window"])] if brief.get("window") else []),
|
|
459
|
+
"noise": str(brief.get("noise") or "normal"),
|
|
460
|
+
"source": "fallback",
|
|
461
|
+
"awaiting_approval": True,
|
|
462
|
+
}
|
|
463
|
+
|
|
464
|
+
def brief_chat(self, payload: dict[str, Any]) -> dict[str, Any]:
|
|
465
|
+
"""A bounded scoping assistant for the audit brief.
|
|
466
|
+
|
|
467
|
+
It answers questions and can suggest brief adjustments, but it has no
|
|
468
|
+
tools and never starts an audit — the audit begins only when the
|
|
469
|
+
operator approves a plan.
|
|
470
|
+
"""
|
|
471
|
+
message = str((payload or {}).get("message") or "").strip()
|
|
472
|
+
brief = (payload or {}).get("brief") or {}
|
|
473
|
+
if not message:
|
|
474
|
+
return {"ok": False, "error": "empty message"}
|
|
475
|
+
try:
|
|
476
|
+
text = asyncio.run(self._brief_chat_async(message, brief))
|
|
477
|
+
except Exception as exc:
|
|
478
|
+
return {"ok": False, "error": f"{type(exc).__name__}: {exc}"}
|
|
479
|
+
return {"ok": True, "text": text}
|
|
480
|
+
|
|
481
|
+
async def _brief_chat_async(self, message: str, brief: dict[str, Any]) -> str:
|
|
482
|
+
from splitagent.agents.prompts import build_planner_chat_prompt
|
|
483
|
+
|
|
484
|
+
project = self.project
|
|
485
|
+
context = SharedContext.create(
|
|
486
|
+
target=project.target.url or "unspecified",
|
|
487
|
+
target_kind=project.target.kind,
|
|
488
|
+
scope=project.target.scope or project.target.effective_hosts(),
|
|
489
|
+
name="planner",
|
|
490
|
+
model=self.global_config.llm.model,
|
|
491
|
+
provider=self.global_config.llm.provider,
|
|
492
|
+
)
|
|
493
|
+
# The *chat* prompt: conversational, not the JSON-only plan prompt.
|
|
494
|
+
system = build_planner_chat_prompt(project, context)
|
|
495
|
+
user = (
|
|
496
|
+
"=== CURRENT BRIEF ===\n"
|
|
497
|
+
+ json.dumps(brief, ensure_ascii=False, indent=2)
|
|
498
|
+
+ "\n\n=== OPERATOR ===\n"
|
|
499
|
+
+ message
|
|
500
|
+
+ "\n\nAnswer briefly in the operator's language. If the brief should "
|
|
501
|
+
"change, list concrete adjustments. You run no tools and start nothing."
|
|
502
|
+
)
|
|
503
|
+
async with LLMClient(self.global_config.llm) as client:
|
|
504
|
+
result, _usage = await asyncio.wait_for(
|
|
505
|
+
client.complete(
|
|
506
|
+
[
|
|
507
|
+
ChatMessage(role="system", content=system),
|
|
508
|
+
ChatMessage(role="user", content=user),
|
|
509
|
+
]
|
|
510
|
+
),
|
|
511
|
+
timeout=PLANNER_TIMEOUT_SECONDS,
|
|
512
|
+
)
|
|
513
|
+
return result.content or "(no answer)"
|
|
514
|
+
|
|
515
|
+
def _build_project(self, options: dict[str, Any]) -> ProjectConfig:
|
|
516
|
+
project = copy.deepcopy(self.project)
|
|
517
|
+
target = options.get("target") or {}
|
|
518
|
+
if target.get("url"):
|
|
519
|
+
project.target.url = str(target["url"])
|
|
520
|
+
if target.get("kind"):
|
|
521
|
+
project.target.kind = str(target["kind"])
|
|
522
|
+
if target.get("scope"):
|
|
523
|
+
project.target.scope = [s.strip() for s in str(target["scope"]).split(",") if s.strip()]
|
|
524
|
+
run = options.get("run") or {}
|
|
525
|
+
if run.get("rounds") is not None:
|
|
526
|
+
try:
|
|
527
|
+
project.run.rounds = max(1, int(run["rounds"]))
|
|
528
|
+
except (TypeError, ValueError):
|
|
529
|
+
pass
|
|
530
|
+
if run.get("max_steps") is not None:
|
|
531
|
+
try:
|
|
532
|
+
project.run.max_steps = max(1, int(run["max_steps"]))
|
|
533
|
+
except (TypeError, ValueError):
|
|
534
|
+
pass
|
|
535
|
+
if run.get("allow_network") is not None:
|
|
536
|
+
project.run.allow_network = bool(run["allow_network"])
|
|
537
|
+
if run.get("sandbox") is not None:
|
|
538
|
+
project.run.sandbox.enabled = bool(run["sandbox"])
|
|
539
|
+
workspace = options.get("workspace")
|
|
540
|
+
if isinstance(workspace, dict):
|
|
541
|
+
if "path" in workspace:
|
|
542
|
+
project.workspace.path = str(workspace.get("path") or "")
|
|
543
|
+
for key in ("allow_install", "allow_external_tools"):
|
|
544
|
+
if workspace.get(key) is not None:
|
|
545
|
+
setattr(project.workspace, key, bool(workspace[key]))
|
|
546
|
+
if workspace.get("max_install_seconds") is not None:
|
|
547
|
+
try:
|
|
548
|
+
project.workspace.max_install_seconds = max(
|
|
549
|
+
30, int(workspace["max_install_seconds"])
|
|
550
|
+
)
|
|
551
|
+
except (TypeError, ValueError):
|
|
552
|
+
pass
|
|
553
|
+
if "instructions" in workspace:
|
|
554
|
+
project.workspace.instructions = str(workspace.get("instructions") or "")
|
|
555
|
+
compaction = options.get("compaction")
|
|
556
|
+
if isinstance(compaction, dict):
|
|
557
|
+
for key in ("auto", "prune"):
|
|
558
|
+
if compaction.get(key) is not None:
|
|
559
|
+
setattr(project.run.compaction, key, bool(compaction[key]))
|
|
560
|
+
for key in ("reserved", "preserve_recent_tokens", "tail_turns"):
|
|
561
|
+
if compaction.get(key) is not None:
|
|
562
|
+
try:
|
|
563
|
+
setattr(project.run.compaction, key, int(compaction[key]))
|
|
564
|
+
except (TypeError, ValueError):
|
|
565
|
+
pass
|
|
566
|
+
auth = options.get("auth") or {}
|
|
567
|
+
if auth:
|
|
568
|
+
if "username" in auth:
|
|
569
|
+
project.auth.username = str(auth.get("username") or "")
|
|
570
|
+
if "password" in auth:
|
|
571
|
+
project.auth.password = str(auth.get("password") or "")
|
|
572
|
+
if "token" in auth:
|
|
573
|
+
project.auth.token = str(auth.get("token") or "")
|
|
574
|
+
if "cookies" in auth:
|
|
575
|
+
project.auth.cookies = str(auth.get("cookies") or "")
|
|
576
|
+
if isinstance(auth.get("headers"), dict):
|
|
577
|
+
project.auth.headers = {str(k): str(v) for k, v in auth["headers"].items() if k}
|
|
578
|
+
brief = options.get("brief")
|
|
579
|
+
if isinstance(brief, dict):
|
|
580
|
+
instructions = _compose_brief_instructions(brief)
|
|
581
|
+
if instructions:
|
|
582
|
+
project.workspace.instructions = (
|
|
583
|
+
(project.workspace.instructions + "\n\n" + instructions).strip()
|
|
584
|
+
if project.workspace.instructions
|
|
585
|
+
else instructions
|
|
586
|
+
)
|
|
587
|
+
noise = str(brief.get("noise") or "").lower()
|
|
588
|
+
if noise == "stealth":
|
|
589
|
+
project.run.tool_concurrency = 1
|
|
590
|
+
project.run.max_duration_minutes = max(project.run.max_duration_minutes, 120)
|
|
591
|
+
elif noise == "aggressive":
|
|
592
|
+
project.run.tool_concurrency = max(project.run.tool_concurrency, 8)
|
|
593
|
+
if options.get("objective"):
|
|
594
|
+
project.name = str(options["objective"])[:60] or project.name
|
|
595
|
+
return project
|
|
596
|
+
|
|
597
|
+
# -- copilot chat ------------------------------------------------------ #
|
|
598
|
+
def is_chatting(self) -> bool:
|
|
599
|
+
return self._chat_thread is not None and self._chat_thread.is_alive()
|
|
600
|
+
|
|
601
|
+
def chat_reset(self) -> dict[str, Any]:
|
|
602
|
+
self._chat_history = []
|
|
603
|
+
self._chat_context = None
|
|
604
|
+
self._emit([{"type": "chat.reset", "data": {}}])
|
|
605
|
+
return {"ok": True}
|
|
606
|
+
|
|
607
|
+
def chat_stop(self) -> dict[str, Any]:
|
|
608
|
+
if self._chat_loop is not None and self._chat_task is not None:
|
|
609
|
+
try:
|
|
610
|
+
self._chat_loop.call_soon_threadsafe(self._chat_task.cancel)
|
|
611
|
+
except RuntimeError:
|
|
612
|
+
pass
|
|
613
|
+
return {"ok": True}
|
|
614
|
+
return {"ok": False, "error": "No chat running."}
|
|
615
|
+
|
|
616
|
+
def chat_send(self, message: str) -> dict[str, Any]:
|
|
617
|
+
text = (message or "").strip()
|
|
618
|
+
if not text:
|
|
619
|
+
return {"ok": False, "error": "empty message"}
|
|
620
|
+
if self.is_chatting():
|
|
621
|
+
return {"ok": False, "error": "The copilot is still answering."}
|
|
622
|
+
self._chat_events = []
|
|
623
|
+
self._chat_done = False
|
|
624
|
+
self._chat_thread = threading.Thread(target=self._chat_entry, args=(text,), daemon=True)
|
|
625
|
+
self._chat_thread.start()
|
|
626
|
+
return {"ok": True}
|
|
627
|
+
|
|
628
|
+
def chat_state(self, since: int = 0) -> dict[str, Any]:
|
|
629
|
+
"""Events since an index, for the UI to poll without gaps.
|
|
630
|
+
|
|
631
|
+
Pushes over the bridge are best-effort; this lets the front-end replay
|
|
632
|
+
whatever it missed, so a reply can never be lost mid-stream.
|
|
633
|
+
"""
|
|
634
|
+
with self._queue_lock:
|
|
635
|
+
events = list(self._chat_events)
|
|
636
|
+
return {
|
|
637
|
+
"ok": True,
|
|
638
|
+
"running": self.is_chatting(),
|
|
639
|
+
"pending": not self._chat_done,
|
|
640
|
+
"events": events[since:],
|
|
641
|
+
"next": len(events),
|
|
642
|
+
}
|
|
643
|
+
|
|
644
|
+
def _chat_entry(self, message: str) -> None:
|
|
645
|
+
try:
|
|
646
|
+
asyncio.run(self._chat_message(message))
|
|
647
|
+
except asyncio.CancelledError:
|
|
648
|
+
self._emit_chat([{"type": "chat.end", "data": {"ok": False, "error": "cancelled"}}])
|
|
649
|
+
self._chat_done = True
|
|
650
|
+
except Exception as exc:
|
|
651
|
+
self._emit_chat(
|
|
652
|
+
[
|
|
653
|
+
{
|
|
654
|
+
"type": "error",
|
|
655
|
+
"agent": "assistant",
|
|
656
|
+
"data": {"text": f"{type(exc).__name__}: {exc}"},
|
|
657
|
+
},
|
|
658
|
+
{"type": "chat.end", "data": {"ok": False, "error": str(exc)}},
|
|
659
|
+
]
|
|
660
|
+
)
|
|
661
|
+
self._chat_done = True
|
|
662
|
+
|
|
663
|
+
async def _chat_message(self, message: str) -> None:
|
|
664
|
+
return await self._chat_async(message)
|
|
665
|
+
|
|
666
|
+
async def _chat_async(self, message: str) -> None:
|
|
667
|
+
if self._chat_context is None:
|
|
668
|
+
target = self.project.target.url or ", ".join(self.project.target.effective_hosts())
|
|
669
|
+
self._chat_context = SharedContext.create(
|
|
670
|
+
target=target or "unspecified",
|
|
671
|
+
target_kind=self.project.target.kind,
|
|
672
|
+
scope=self.project.target.scope or self.project.target.effective_hosts(),
|
|
673
|
+
name="copilot",
|
|
674
|
+
model=self.global_config.llm.model,
|
|
675
|
+
provider=self.global_config.llm.provider,
|
|
676
|
+
)
|
|
677
|
+
bus = EventBus()
|
|
678
|
+
# Chat events feed only the polled buffer (single channel, no bridge).
|
|
679
|
+
bus.subscribe(lambda event: self._emit_chat([_event_dict(event)]))
|
|
680
|
+
self._emit_chat([{"type": "chat.start", "data": {"text": message}}])
|
|
681
|
+
|
|
682
|
+
self._chat_loop = asyncio.get_running_loop()
|
|
683
|
+
async with LLMClient(self.global_config.llm) as client:
|
|
684
|
+
agent = ChatAgent(
|
|
685
|
+
client=client,
|
|
686
|
+
context=self._chat_context,
|
|
687
|
+
project=self.project,
|
|
688
|
+
bus=bus,
|
|
689
|
+
settings={
|
|
690
|
+
"sandbox": None,
|
|
691
|
+
"auth_headers": self.project.auth.as_headers(),
|
|
692
|
+
},
|
|
693
|
+
history=self._chat_history,
|
|
694
|
+
)
|
|
695
|
+
self._chat_task = asyncio.ensure_future(agent.run(message))
|
|
696
|
+
while not self._chat_task.done():
|
|
697
|
+
await asyncio.sleep(0.04)
|
|
698
|
+
result = self._chat_task.result()
|
|
699
|
+
self._chat_history = list(agent.history)
|
|
700
|
+
self._emit_chat([{"type": "chat.end", "data": {"ok": True, "text": result.text}}])
|
|
701
|
+
self._chat_done = True
|
|
702
|
+
|
|
703
|
+
# -- toolbox ----------------------------------------------------------- #
|
|
704
|
+
def _toolbox(self) -> Toolbox:
|
|
705
|
+
workspace = workspace_for(self.project, base=Path.cwd())
|
|
706
|
+
return Toolbox(self.project.run.execution, workspace.root)
|
|
707
|
+
|
|
708
|
+
def toolbox_status(self, probe: bool = False) -> dict[str, Any]:
|
|
709
|
+
toolbox = self._toolbox()
|
|
710
|
+
status = toolbox.detect(probe=probe)
|
|
711
|
+
effective = resolve_mode(self.project.run.execution, status)
|
|
712
|
+
data = status.to_dict()
|
|
713
|
+
data["effective_mode"] = effective
|
|
714
|
+
data["installed"] = self.project.run.execution.installed
|
|
715
|
+
data["build_command"] = " ".join(build_command(self.project.run.execution.edition))
|
|
716
|
+
data["estimated_size"] = estimated_size(self.project.run.execution.edition)
|
|
717
|
+
data["workspace"] = str(toolbox.workspace_root)
|
|
718
|
+
return {"ok": True, "status": data}
|
|
719
|
+
|
|
720
|
+
def is_toolbox_busy(self) -> bool:
|
|
721
|
+
return self._toolbox_thread is not None and self._toolbox_thread.is_alive()
|
|
722
|
+
|
|
723
|
+
def toolbox_setup(self, edition: str = "", build: bool = True) -> dict[str, Any]:
|
|
724
|
+
"""Kick off the one-time setup in the background.
|
|
725
|
+
|
|
726
|
+
A Docker build takes minutes, so this must never block the UI thread:
|
|
727
|
+
the webview would freeze and the progress events would never flush.
|
|
728
|
+
Progress and the final result arrive through the event bus.
|
|
729
|
+
"""
|
|
730
|
+
if self.is_toolbox_busy():
|
|
731
|
+
return {"ok": False, "error": "setup already running", "async": True}
|
|
732
|
+
config = self.project.run.execution
|
|
733
|
+
if edition in ("standard", "kali"):
|
|
734
|
+
config.edition = edition
|
|
735
|
+
path = project_config_path()
|
|
736
|
+
if path.exists():
|
|
737
|
+
save_project_config(self.project, path)
|
|
738
|
+
self._toolbox_thread = threading.Thread(
|
|
739
|
+
target=self._toolbox_setup_worker, args=(build,), daemon=True
|
|
740
|
+
)
|
|
741
|
+
self._toolbox_thread.start()
|
|
742
|
+
return {"ok": True, "async": True, "status": self.toolbox_status()["status"]}
|
|
743
|
+
|
|
744
|
+
def toolbox_probe(self) -> dict[str, Any]:
|
|
745
|
+
"""Slow path: ask the container which tools it has (used sparingly)."""
|
|
746
|
+
return self.toolbox_status(probe=True)
|
|
747
|
+
|
|
748
|
+
def _toolbox_setup_worker(self, build: bool) -> None:
|
|
749
|
+
config = self.project.run.execution
|
|
750
|
+
toolbox = self._toolbox()
|
|
751
|
+
|
|
752
|
+
def progress(stage: str, text: str) -> None:
|
|
753
|
+
self._emit([{"type": "toolbox.progress", "data": {"stage": stage, "text": text}}])
|
|
754
|
+
self._flush(force=True)
|
|
755
|
+
|
|
756
|
+
try:
|
|
757
|
+
status = toolbox.detect()
|
|
758
|
+
if not status.docker_cli:
|
|
759
|
+
self._toolbox_fail("Docker is not installed. Install Docker Desktop and retry.")
|
|
760
|
+
return
|
|
761
|
+
|
|
762
|
+
if not status.daemon:
|
|
763
|
+
if not config.auto_start:
|
|
764
|
+
self._toolbox_fail("Docker is not running.")
|
|
765
|
+
return
|
|
766
|
+
progress("daemon", "Starting Docker…")
|
|
767
|
+
started, message = toolbox.start_daemon()
|
|
768
|
+
if not started:
|
|
769
|
+
self._toolbox_fail(message)
|
|
770
|
+
return
|
|
771
|
+
|
|
772
|
+
if build and not toolbox.detect().image:
|
|
773
|
+
progress(
|
|
774
|
+
"build",
|
|
775
|
+
f"Building the {config.edition} toolbox image ({estimated_size(config.edition)})…",
|
|
776
|
+
)
|
|
777
|
+
built = toolbox.build()
|
|
778
|
+
if not built["ok"]:
|
|
779
|
+
self._toolbox_fail(
|
|
780
|
+
built.get("error", "build failed"),
|
|
781
|
+
output=built.get("output", ""),
|
|
782
|
+
)
|
|
783
|
+
return
|
|
784
|
+
|
|
785
|
+
progress("start", "Starting the toolbox…")
|
|
786
|
+
result = toolbox.up(build_if_missing=False)
|
|
787
|
+
if not result.get("ok"):
|
|
788
|
+
self._toolbox_fail(result.get("error", "could not start the toolbox"))
|
|
789
|
+
return
|
|
790
|
+
|
|
791
|
+
config.installed = True
|
|
792
|
+
# ``config`` lives on the project, so persist the project too.
|
|
793
|
+
save_global_config(self.global_config)
|
|
794
|
+
save_project_config(self.project, project_config_path())
|
|
795
|
+
status = self.toolbox_status()["status"]
|
|
796
|
+
self._emit([{"type": "toolbox.ready", "data": status}])
|
|
797
|
+
self._flush(force=True)
|
|
798
|
+
except Exception as exc:
|
|
799
|
+
self._toolbox_fail(f"{type(exc).__name__}: {exc}")
|
|
800
|
+
|
|
801
|
+
def _toolbox_fail(self, error: str, output: str = "") -> None:
|
|
802
|
+
self.project.run.execution.installed = False
|
|
803
|
+
self._emit(
|
|
804
|
+
[
|
|
805
|
+
{
|
|
806
|
+
"type": "toolbox.failed",
|
|
807
|
+
"data": {"error": error, "output": output[-4000:]},
|
|
808
|
+
}
|
|
809
|
+
]
|
|
810
|
+
)
|
|
811
|
+
self._flush(force=True)
|
|
812
|
+
|
|
813
|
+
def _toolbox_action_worker(self, action: str) -> None:
|
|
814
|
+
toolbox = self._toolbox()
|
|
815
|
+
try:
|
|
816
|
+
if action == "start":
|
|
817
|
+
result = toolbox.up()
|
|
818
|
+
elif action == "reset":
|
|
819
|
+
result = toolbox.reset()
|
|
820
|
+
elif action == "rebuild":
|
|
821
|
+
toolbox.remove_image()
|
|
822
|
+
result = toolbox.up()
|
|
823
|
+
else:
|
|
824
|
+
result = {"ok": False, "error": f"unknown action '{action}'"}
|
|
825
|
+
if not result.get("ok"):
|
|
826
|
+
self._toolbox_fail(result.get("error", f"{action} failed"))
|
|
827
|
+
return
|
|
828
|
+
self._emit([{"type": "toolbox.ready", "data": self.toolbox_status()["status"]}])
|
|
829
|
+
self._flush(force=True)
|
|
830
|
+
except Exception as exc:
|
|
831
|
+
self._toolbox_fail(f"{type(exc).__name__}: {exc}")
|
|
832
|
+
|
|
833
|
+
def toolbox_action(self, action: str, edition: str = "") -> dict[str, Any]:
|
|
834
|
+
toolbox = self._toolbox()
|
|
835
|
+
# Anything that can take more than a moment runs in the background so
|
|
836
|
+
# the window never freezes and no console windows pile up.
|
|
837
|
+
if action in ("start", "reset", "rebuild"):
|
|
838
|
+
if self.is_toolbox_busy():
|
|
839
|
+
return {"ok": False, "error": "another toolbox task is running", "async": True}
|
|
840
|
+
self._toolbox_thread = threading.Thread(
|
|
841
|
+
target=self._toolbox_action_worker, args=(action,), daemon=True
|
|
842
|
+
)
|
|
843
|
+
self._toolbox_thread.start()
|
|
844
|
+
return {"ok": True, "async": True}
|
|
845
|
+
if action == "stop":
|
|
846
|
+
toolbox.down()
|
|
847
|
+
return {"ok": True}
|
|
848
|
+
if action == "setup":
|
|
849
|
+
return self.toolbox_setup(edition=edition)
|
|
850
|
+
if action == "set_edition":
|
|
851
|
+
if edition in ("standard", "kali"):
|
|
852
|
+
self.project.run.execution.edition = edition
|
|
853
|
+
path = project_config_path()
|
|
854
|
+
if path.exists():
|
|
855
|
+
save_project_config(self.project, path)
|
|
856
|
+
return {"ok": True, "edition": self.project.run.execution.edition}
|
|
857
|
+
if action == "status":
|
|
858
|
+
return self.toolbox_status()
|
|
859
|
+
return {"ok": False, "error": f"unknown action '{action}'"}
|
|
860
|
+
|
|
861
|
+
def toolbox_install(self, manager: str, package: str) -> dict[str, Any]:
|
|
862
|
+
toolbox = self._toolbox()
|
|
863
|
+
status = toolbox.detect()
|
|
864
|
+
if not status.running:
|
|
865
|
+
return {"ok": False, "error": "the toolbox is not running"}
|
|
866
|
+
return toolbox.install(manager, package)
|
|
867
|
+
|
|
868
|
+
# -- models & providers ------------------------------------------------ #
|
|
869
|
+
def list_models(self) -> dict[str, Any]:
|
|
870
|
+
"""Aggregate the catalogues of every connected provider."""
|
|
871
|
+
config = self.global_config
|
|
872
|
+
entries: list[dict[str, Any]] = []
|
|
873
|
+
for provider_id, provider in config.providers.items():
|
|
874
|
+
catalogue = dict(provider.models)
|
|
875
|
+
if not catalogue:
|
|
876
|
+
preset = PROVIDER_PRESETS.get(provider_id, {})
|
|
877
|
+
default = preset.get("model", config.llm.model)
|
|
878
|
+
if default:
|
|
879
|
+
catalogue = {default: default}
|
|
880
|
+
for model_id, name in catalogue.items():
|
|
881
|
+
entries.append(
|
|
882
|
+
{
|
|
883
|
+
"provider": provider_id,
|
|
884
|
+
"provider_name": provider_id,
|
|
885
|
+
"id": model_id,
|
|
886
|
+
"name": name or model_id,
|
|
887
|
+
"source": "connected",
|
|
888
|
+
"visible": provider.visible(model_id),
|
|
889
|
+
"free": model_id.endswith("-free"),
|
|
890
|
+
}
|
|
891
|
+
)
|
|
892
|
+
return {
|
|
893
|
+
"ok": True,
|
|
894
|
+
"models": entries,
|
|
895
|
+
"current": {"provider": config.llm.provider, "model": config.llm.model},
|
|
896
|
+
"connected": list(config.providers),
|
|
897
|
+
}
|
|
898
|
+
|
|
899
|
+
def list_providers(self) -> dict[str, Any]:
|
|
900
|
+
config = self.global_config
|
|
901
|
+
connected = []
|
|
902
|
+
for provider_id, provider in config.providers.items():
|
|
903
|
+
connected.append(
|
|
904
|
+
{
|
|
905
|
+
"id": provider_id,
|
|
906
|
+
"name": provider_id,
|
|
907
|
+
"base_url": provider.base_url,
|
|
908
|
+
"protocol": provider.protocol,
|
|
909
|
+
"enabled": provider.enabled,
|
|
910
|
+
"models": len(provider.models),
|
|
911
|
+
"hidden": len(provider.hidden),
|
|
912
|
+
"has_key": bool(provider.api_key),
|
|
913
|
+
"active": provider_id == config.llm.provider,
|
|
914
|
+
}
|
|
915
|
+
)
|
|
916
|
+
connected.sort(key=lambda item: (not item["active"], item["id"]))
|
|
917
|
+
popular = [
|
|
918
|
+
{
|
|
919
|
+
"id": name,
|
|
920
|
+
"name": name,
|
|
921
|
+
"base_url": preset["base_url"],
|
|
922
|
+
"protocol": preset.get("protocol", "openai"),
|
|
923
|
+
"default_model": preset["model"],
|
|
924
|
+
}
|
|
925
|
+
for name, preset in PROVIDER_PRESETS.items()
|
|
926
|
+
if name not in config.providers
|
|
927
|
+
]
|
|
928
|
+
return {"ok": True, "connected": connected, "popular": popular}
|
|
929
|
+
|
|
930
|
+
def _provider_settings(
|
|
931
|
+
self, provider_id: str, base_url: str, protocol: str, api_key: str, model: str
|
|
932
|
+
) -> LLMSettings:
|
|
933
|
+
settings = LLMSettings(**vars(self.global_config.llm))
|
|
934
|
+
settings.provider = provider_id
|
|
935
|
+
settings.base_url = base_url or settings.base_url
|
|
936
|
+
settings.protocol = protocol or "openai"
|
|
937
|
+
settings.api_key = api_key
|
|
938
|
+
if model:
|
|
939
|
+
settings.model = model
|
|
940
|
+
return settings
|
|
941
|
+
|
|
942
|
+
def connect_provider(self, payload: dict[str, Any]) -> dict[str, Any]:
|
|
943
|
+
config = self.global_config
|
|
944
|
+
provider_id = str(payload.get("provider") or "").strip()
|
|
945
|
+
if not provider_id:
|
|
946
|
+
return {"ok": False, "error": "provider is required"}
|
|
947
|
+
preset = PROVIDER_PRESETS.get(provider_id, {})
|
|
948
|
+
base_url = str(payload.get("base_url") or preset.get("base_url") or "").strip()
|
|
949
|
+
protocol = str(payload.get("protocol") or preset.get("protocol") or "openai")
|
|
950
|
+
api_key = str(payload.get("api_key") or "").strip()
|
|
951
|
+
if not base_url:
|
|
952
|
+
return {"ok": False, "error": "base URL is required"}
|
|
953
|
+
|
|
954
|
+
provider = config.providers.get(provider_id) or ProviderConfig(id=provider_id)
|
|
955
|
+
provider.base_url = base_url
|
|
956
|
+
provider.protocol = protocol
|
|
957
|
+
if api_key:
|
|
958
|
+
provider.api_key = api_key
|
|
959
|
+
provider.enabled = True
|
|
960
|
+
|
|
961
|
+
warning = ""
|
|
962
|
+
fetched: list[dict[str, str]] = []
|
|
963
|
+
settings = self._provider_settings(
|
|
964
|
+
provider_id, base_url, protocol, provider.api_key, preset.get("model", "")
|
|
965
|
+
)
|
|
966
|
+
try:
|
|
967
|
+
fetched = asyncio.run(_fetch_models_async(settings))
|
|
968
|
+
except Exception as exc:
|
|
969
|
+
warning = f"could not list models: {exc}"
|
|
970
|
+
|
|
971
|
+
for model in fetched:
|
|
972
|
+
provider.models.setdefault(model["id"], model["name"])
|
|
973
|
+
if not provider.models:
|
|
974
|
+
default = preset.get("model") or config.llm.model
|
|
975
|
+
if default:
|
|
976
|
+
provider.models[default] = default
|
|
977
|
+
|
|
978
|
+
config.providers[provider_id] = provider
|
|
979
|
+
if payload.get("activate") or not config.llm.model or provider_id == config.llm.provider:
|
|
980
|
+
self.activate_model(
|
|
981
|
+
provider_id, config.llm.model if provider_id == config.llm.provider else ""
|
|
982
|
+
)
|
|
983
|
+
save_global_config(config)
|
|
984
|
+
return {
|
|
985
|
+
"ok": True,
|
|
986
|
+
"provider": provider_id,
|
|
987
|
+
"models": len(provider.models),
|
|
988
|
+
"warning": warning,
|
|
989
|
+
}
|
|
990
|
+
|
|
991
|
+
def _fetch_models(self, settings: LLMSettings) -> list[dict[str, str]]:
|
|
992
|
+
return asyncio.run(_fetch_models_async(settings))
|
|
993
|
+
|
|
994
|
+
def refresh_provider(self, provider_id: str) -> dict[str, Any]:
|
|
995
|
+
config = self.global_config
|
|
996
|
+
provider = config.providers.get(provider_id)
|
|
997
|
+
if provider is None:
|
|
998
|
+
return {"ok": False, "error": f"provider '{provider_id}' is not connected"}
|
|
999
|
+
settings = self._provider_settings(
|
|
1000
|
+
provider_id, provider.base_url, provider.protocol, provider.api_key, ""
|
|
1001
|
+
)
|
|
1002
|
+
try:
|
|
1003
|
+
fetched = self._fetch_models(settings)
|
|
1004
|
+
except Exception as exc:
|
|
1005
|
+
return {"ok": False, "error": str(exc)}
|
|
1006
|
+
for model in fetched:
|
|
1007
|
+
provider.models.setdefault(model["id"], model["name"])
|
|
1008
|
+
save_global_config(config)
|
|
1009
|
+
return {"ok": True, "models": len(provider.models)}
|
|
1010
|
+
|
|
1011
|
+
def disconnect_provider(self, provider_id: str) -> dict[str, Any]:
|
|
1012
|
+
config = self.global_config
|
|
1013
|
+
config.providers.pop(provider_id, None)
|
|
1014
|
+
if config.llm.provider == provider_id:
|
|
1015
|
+
remaining = next(iter(config.providers), None)
|
|
1016
|
+
if remaining:
|
|
1017
|
+
self.activate_model(remaining, "")
|
|
1018
|
+
else:
|
|
1019
|
+
config.llm.api_key = ""
|
|
1020
|
+
save_global_config(config)
|
|
1021
|
+
return {"ok": True}
|
|
1022
|
+
|
|
1023
|
+
def activate_model(self, provider_id: str, model: str = "") -> dict[str, Any]:
|
|
1024
|
+
config = self.global_config
|
|
1025
|
+
provider = config.providers.get(provider_id)
|
|
1026
|
+
chosen = model
|
|
1027
|
+
if not chosen and provider is not None:
|
|
1028
|
+
chosen = next(iter(provider.models), "")
|
|
1029
|
+
if not chosen:
|
|
1030
|
+
preset = PROVIDER_PRESETS.get(provider_id, {})
|
|
1031
|
+
chosen = preset.get("model", config.llm.model)
|
|
1032
|
+
settings = settings_for_provider(config, provider_id, chosen)
|
|
1033
|
+
config.llm.provider = settings.provider
|
|
1034
|
+
config.llm.base_url = settings.base_url
|
|
1035
|
+
config.llm.protocol = settings.protocol
|
|
1036
|
+
config.llm.model = settings.model
|
|
1037
|
+
if settings.api_key:
|
|
1038
|
+
config.llm.api_key = settings.api_key
|
|
1039
|
+
save_global_config(config)
|
|
1040
|
+
return {"ok": True, "provider": provider_id, "model": config.llm.model}
|
|
1041
|
+
|
|
1042
|
+
def set_model_visibility(self, provider_id: str, model: str, visible: bool) -> dict[str, Any]:
|
|
1043
|
+
config = self.global_config
|
|
1044
|
+
provider = config.providers.get(provider_id)
|
|
1045
|
+
if provider is None:
|
|
1046
|
+
return {"ok": False, "error": "provider not connected"}
|
|
1047
|
+
if visible:
|
|
1048
|
+
provider.hidden = [m for m in provider.hidden if m != model]
|
|
1049
|
+
elif model not in provider.hidden:
|
|
1050
|
+
provider.hidden.append(model)
|
|
1051
|
+
save_global_config(config)
|
|
1052
|
+
return {"ok": True}
|
|
1053
|
+
|
|
1054
|
+
def set_provider_visibility(self, provider_id: str, visible: bool) -> dict[str, Any]:
|
|
1055
|
+
config = self.global_config
|
|
1056
|
+
provider = config.providers.get(provider_id)
|
|
1057
|
+
if provider is None:
|
|
1058
|
+
return {"ok": False, "error": "provider not connected"}
|
|
1059
|
+
provider.enabled = visible
|
|
1060
|
+
save_global_config(config)
|
|
1061
|
+
return {"ok": True}
|
|
1062
|
+
|
|
1063
|
+
def add_custom_model(self, provider_id: str, model: str, name: str = "") -> dict[str, Any]:
|
|
1064
|
+
config = self.global_config
|
|
1065
|
+
model = (model or "").strip()
|
|
1066
|
+
if not model:
|
|
1067
|
+
return {"ok": False, "error": "model id is required"}
|
|
1068
|
+
provider = ensure_provider(config, provider_id)
|
|
1069
|
+
provider.models[model] = (name or model).strip()
|
|
1070
|
+
provider.hidden = [m for m in provider.hidden if m != model]
|
|
1071
|
+
save_global_config(config)
|
|
1072
|
+
return {"ok": True, "provider": provider_id, "model": model}
|
|
1073
|
+
|
|
1074
|
+
def remove_model(self, provider_id: str, model: str) -> dict[str, Any]:
|
|
1075
|
+
config = self.global_config
|
|
1076
|
+
provider = config.providers.get(provider_id)
|
|
1077
|
+
if provider is None:
|
|
1078
|
+
return {"ok": False, "error": "provider not connected"}
|
|
1079
|
+
provider.models.pop(model, None)
|
|
1080
|
+
provider.hidden = [m for m in provider.hidden if m != model]
|
|
1081
|
+
save_global_config(config)
|
|
1082
|
+
return {"ok": True}
|
|
1083
|
+
|
|
1084
|
+
# -- sessions / reports ------------------------------------------------ #
|
|
1085
|
+
def list_sessions(self) -> list[dict[str, Any]]:
|
|
1086
|
+
directory = sessions_dir()
|
|
1087
|
+
if not directory.exists():
|
|
1088
|
+
return []
|
|
1089
|
+
entries: list[dict[str, Any]] = []
|
|
1090
|
+
for path in directory.glob("*.session.enc"):
|
|
1091
|
+
stat = path.stat()
|
|
1092
|
+
entries.append(
|
|
1093
|
+
{
|
|
1094
|
+
"id": path.name.replace(".session.enc", ""),
|
|
1095
|
+
"modified": datetime.fromtimestamp(stat.st_mtime, tz=timezone.utc).isoformat(
|
|
1096
|
+
timespec="seconds"
|
|
1097
|
+
),
|
|
1098
|
+
"size": stat.st_size,
|
|
1099
|
+
}
|
|
1100
|
+
)
|
|
1101
|
+
entries.sort(key=lambda item: item["modified"], reverse=True)
|
|
1102
|
+
return entries[:40]
|
|
1103
|
+
|
|
1104
|
+
def load_session(self, session_id: str) -> dict[str, Any]:
|
|
1105
|
+
try:
|
|
1106
|
+
context = SharedContext.load(session_id)
|
|
1107
|
+
except Exception as exc:
|
|
1108
|
+
return {"ok": False, "error": str(exc)}
|
|
1109
|
+
self.session = context
|
|
1110
|
+
return {
|
|
1111
|
+
"ok": True,
|
|
1112
|
+
"state": context.state.to_dict(),
|
|
1113
|
+
"summary": context.summary_dict(),
|
|
1114
|
+
}
|
|
1115
|
+
|
|
1116
|
+
def export_report(self, formats: list[str] | None = None) -> dict[str, Any]:
|
|
1117
|
+
if self.session is None:
|
|
1118
|
+
return {"ok": False, "error": "No session to export."}
|
|
1119
|
+
from splitagent.report.generator import write_reports
|
|
1120
|
+
|
|
1121
|
+
report_settings = self.project.report
|
|
1122
|
+
if formats:
|
|
1123
|
+
report_settings.formats = list(formats)
|
|
1124
|
+
try:
|
|
1125
|
+
paths = write_reports(
|
|
1126
|
+
self.session.state, report_settings, Path(report_settings.output_dir)
|
|
1127
|
+
)
|
|
1128
|
+
except (OSError, ValueError) as exc:
|
|
1129
|
+
return {"ok": False, "error": f"could not write report: {exc}"}
|
|
1130
|
+
self.report_paths = [str(p.resolve()) for p in paths]
|
|
1131
|
+
return {"ok": True, "reports": self.report_paths}
|