Skip to content

Commit f2f0cff

Browse files
committed
fix: address sandbox provider review feedback
1 parent 211de97 commit f2f0cff

15 files changed

Lines changed: 213 additions & 52 deletions

astrbot/core/astr_agent_tool_exec.py

Lines changed: 8 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -20,6 +20,9 @@
2020
from astrbot.core.astr_main_agent_resources import (
2121
BACKGROUND_TASK_RESULT_WOKE_SYSTEM_PROMPT,
2222
)
23+
from astrbot.core.computer.sandbox_tool_binding import (
24+
resolve_sandbox_provider_bindings,
25+
)
2326
from astrbot.core.cron.events import CronMessageEvent
2427
from astrbot.core.message.components import Image
2528
from astrbot.core.message.message_event_result import (
@@ -237,14 +240,11 @@ def _get_runtime_computer_tools(
237240
edit_tool.name: edit_tool,
238241
grep_tool.name: grep_tool,
239242
}
240-
from astrbot.core.computer.computer_client import get_sandbox_provider_info
241-
242-
provider_info = get_sandbox_provider_info(booter)
243-
if provider_info:
244-
for tool_name in provider_info.get("tool_names", []):
245-
provider_tool = tool_mgr.get_func(tool_name)
246-
if provider_tool and getattr(provider_tool, "active", True):
247-
tools[provider_tool.name] = provider_tool
243+
_provider_info, provider_tools = resolve_sandbox_provider_bindings(
244+
booter, tool_mgr
245+
)
246+
for provider_tool in provider_tools:
247+
tools[provider_tool.name] = provider_tool
248248
return tools
249249
if runtime == "local":
250250
shell_tool = tool_mgr.get_builtin_tool(ExecuteShellTool)

astrbot/core/astr_main_agent.py

Lines changed: 9 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -1002,16 +1002,18 @@ def _apply_sandbox_tools(
10021002
req.func_tool.add_tool(tool_mgr.get_builtin_tool(FileWriteTool))
10031003
req.func_tool.add_tool(tool_mgr.get_builtin_tool(FileEditTool))
10041004
req.func_tool.add_tool(tool_mgr.get_builtin_tool(GrepTool))
1005-
from astrbot.core.computer.computer_client import get_sandbox_provider_info
1005+
from astrbot.core.computer.sandbox_tool_binding import (
1006+
resolve_sandbox_provider_bindings,
1007+
)
10061008

1007-
provider_info = get_sandbox_provider_info(booter)
1008-
if provider_info:
1009-
for tool_name in provider_info.get("tool_names", []):
1010-
tool = tool_mgr.get_func(tool_name)
1011-
if tool and getattr(tool, "active", True):
1012-
req.func_tool.add_tool(tool)
1009+
provider_info, provider_tools = resolve_sandbox_provider_bindings(booter, tool_mgr)
1010+
for tool in provider_tools:
1011+
req.func_tool.add_tool(tool)
10131012

10141013
req.system_prompt = f"{req.system_prompt or ''}\n{SANDBOX_MODE_PROMPT}\n"
1014+
provider_system_prompt = (provider_info or {}).get("system_prompt", "").strip()
1015+
if provider_system_prompt:
1016+
req.system_prompt = f"{req.system_prompt or ''}\n{provider_system_prompt}\n"
10151017

10161018

10171019
def _proactive_cron_job_tools(req: ProviderRequest, plugin_context: Context) -> None:

astrbot/core/computer/computer_client.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -29,6 +29,7 @@ def _sandbox_provider_info(provider_id: str, provider: SandboxProvider) -> dict:
2929
"provider_id": provider_id,
3030
"capabilities": sorted(getattr(provider, "capabilities", set())),
3131
"tool_names": sorted(getattr(provider, "tool_names", set())),
32+
"system_prompt": str(getattr(provider, "system_prompt", "") or ""),
3233
}
3334

3435

astrbot/core/computer/sandbox_manager.py

Lines changed: 51 additions & 35 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,7 @@
11
from __future__ import annotations
22

33
import asyncio
4+
import inspect
45
import time
56
import uuid
67
from dataclasses import dataclass
@@ -39,6 +40,12 @@ def save_registry(self) -> None:
3940
except Exception as exc:
4041
logger.warning("[Computer] Failed to save sandbox registry: %s", exc)
4142

43+
async def save_registry_async(self) -> None:
44+
try:
45+
await self.registry.save_async()
46+
except Exception as exc:
47+
logger.warning("[Computer] Failed to save sandbox registry: %s", exc)
48+
4249
def _sandbox_boot_lock(self, sandbox_id: str) -> asyncio.Lock:
4350
lock = self.boot_locks.get(sandbox_id)
4451
if lock is None:
@@ -101,7 +108,10 @@ async def booter_available(self, booter: ComputerBooter) -> bool:
101108
available = getattr(booter, "available", None)
102109
if available is None:
103110
return True
104-
return await available()
111+
result = available() if callable(available) else available
112+
if inspect.isawaitable(result):
113+
result = await result
114+
return bool(result)
105115

106116
def acquire_lease(
107117
self, sandbox_id: str, session_id: str, *, ttl: float | None = None
@@ -137,7 +147,7 @@ def sandbox_controlled_by_other_session(
137147
return False
138148
return bool(lease_expires_at and lease_expires_at > time.time())
139149

140-
def _upsert_new_sandbox_record(
150+
async def _upsert_new_sandbox_record(
141151
self, context: Context, session_id: str, provider_id: str, create_config: dict
142152
) -> str:
143153
provider = self.get_provider(provider_id)
@@ -152,7 +162,7 @@ def _upsert_new_sandbox_record(
152162
connect_info=provider.build_connect_info(sandbox_id, create_config),
153163
)
154164
)
155-
self.save_registry()
165+
await self.save_registry_async()
156166
return sandbox_id
157167

158168
async def get_or_create_booter(
@@ -168,22 +178,26 @@ async def get_or_create_booter(
168178
current_record is None or current_record.get("provider") != provider_id
169179
):
170180
self.registry.set_current_sandbox_id(session_id, None)
171-
self.save_registry()
181+
await self.save_registry_async()
182+
current_sandbox_id = None
183+
current_record = None
172184
if (
173185
current_sandbox_id
174186
and current_record
175187
and current_record.get("provider") == provider_id
176188
and current_sandbox_id in self.session_booter
177189
):
178190
if not self.acquire_lease(current_sandbox_id, session_id):
179-
raise RuntimeError(f"Sandbox {current_sandbox_id} is busy")
180-
booter = self.session_booter[current_sandbox_id]
181-
if await self.booter_available(booter):
182-
self.registry.touch_sandbox(current_sandbox_id)
183-
self.save_registry()
184-
self.schedule_idle_cleanup(current_sandbox_id, idle_timeout)
185-
return booter
186-
self.session_booter.pop(current_sandbox_id, None)
191+
self.registry.set_current_sandbox_id(session_id, None)
192+
await self.save_registry_async()
193+
else:
194+
booter = self.session_booter[current_sandbox_id]
195+
if await self.booter_available(booter):
196+
self.registry.touch_sandbox(current_sandbox_id)
197+
await self.save_registry_async()
198+
self.schedule_idle_cleanup(current_sandbox_id, idle_timeout)
199+
return booter
200+
self.session_booter.pop(current_sandbox_id, None)
187201

188202
created_target_record = False
189203
target_sandbox_id = self.get_default_sandbox_id(provider_id)
@@ -204,10 +218,10 @@ async def get_or_create_booter(
204218
)
205219
)
206220
self.registry.set_default_sandbox_id(record["sandbox_id"])
207-
self.save_registry()
221+
await self.save_registry_async()
208222

209223
if self.sandbox_controlled_by_other_session(target_sandbox_id, session_id):
210-
target_sandbox_id = self._upsert_new_sandbox_record(
224+
target_sandbox_id = await self._upsert_new_sandbox_record(
211225
context, session_id, provider_id, create_config
212226
)
213227
created_target_record = True
@@ -217,7 +231,7 @@ async def get_or_create_booter(
217231
if target_sandbox_id in self.session_booter and not self.acquire_lease(
218232
target_sandbox_id, session_id
219233
):
220-
target_sandbox_id = self._upsert_new_sandbox_record(
234+
target_sandbox_id = await self._upsert_new_sandbox_record(
221235
context, session_id, provider_id, create_config
222236
)
223237
created_target_record = True
@@ -230,10 +244,10 @@ async def get_or_create_booter(
230244
self.session_booter.pop(target_sandbox_id, None)
231245
self.registry.release_lease(target_sandbox_id)
232246
self.registry.update_sandbox_status(target_sandbox_id, "unknown")
233-
self.save_registry()
247+
await self.save_registry_async()
234248

235249
if not self.acquire_lease(target_sandbox_id, session_id):
236-
target_sandbox_id = self._upsert_new_sandbox_record(
250+
target_sandbox_id = await self._upsert_new_sandbox_record(
237251
context, session_id, provider_id, create_config
238252
)
239253
created_target_record = True
@@ -252,7 +266,7 @@ async def get_or_create_booter(
252266
target_sandbox_id, "unknown"
253267
)
254268
self.drop_boot_lock(target_sandbox_id)
255-
self.save_registry()
269+
await self.save_registry_async()
256270
raise
257271
setattr(client, "sandbox_id", target_sandbox_id)
258272
self.session_booter[target_sandbox_id] = client
@@ -262,7 +276,7 @@ async def get_or_create_booter(
262276
self.registry.touch_sandbox(target_sandbox_id)
263277
self.registry.update_sandbox_status(target_sandbox_id, "running")
264278
self.registry.set_current_sandbox_id(session_id, target_sandbox_id)
265-
self.save_registry()
279+
await self.save_registry_async()
266280
self.schedule_idle_cleanup(target_sandbox_id, idle_timeout)
267281
return self.session_booter[target_sandbox_id]
268282

@@ -295,13 +309,13 @@ async def create_sandbox_uncontrolled(
295309
except Exception:
296310
self.registry.delete_sandbox(sandbox_id)
297311
self.drop_boot_lock(sandbox_id)
298-
self.save_registry()
312+
await self.save_registry_async()
299313
raise
300314
setattr(client, "sandbox_id", sandbox_id)
301315
self.session_booter[sandbox_id] = client
302316
self.registry.touch_sandbox(sandbox_id)
303317
self.registry.update_sandbox_status(sandbox_id, "running")
304-
self.save_registry()
318+
await self.save_registry_async()
305319
self.schedule_idle_cleanup(sandbox_id, idle_timeout)
306320
return self.registry.get_sandbox(sandbox_id) or record
307321

@@ -319,7 +333,7 @@ async def create_sandbox(
319333
if not self.acquire_lease(sandbox_id, session_id):
320334
raise RuntimeError(f"Sandbox {sandbox_id} is busy")
321335
self.registry.set_current_sandbox_id(session_id, sandbox_id)
322-
self.save_registry()
336+
await self.save_registry_async()
323337
return self.registry.get_sandbox(sandbox_id) or sandbox
324338

325339
def list_sandboxes(self) -> list[dict]:
@@ -395,13 +409,15 @@ async def switch_current_sandbox_checked(
395409
if not await self.booter_available(booter):
396410
self.session_booter.pop(sandbox_id, None)
397411
self.registry.update_sandbox_status(sandbox_id, "unknown")
398-
self.save_registry()
412+
await self.save_registry_async()
399413
raise RuntimeError(f"Sandbox {sandbox_id} is not running")
400414
if not self.acquire_lease(sandbox_id, session_id):
401415
raise RuntimeError(f"Sandbox {sandbox_id} is busy")
402-
return self._set_current_sandbox_after_lease(session_id, sandbox_id, record)
416+
return await self._set_current_sandbox_after_lease(
417+
session_id, sandbox_id, record
418+
)
403419

404-
def _set_current_sandbox_after_lease(
420+
async def _set_current_sandbox_after_lease(
405421
self, session_id: str, sandbox_id: str, record: dict
406422
) -> dict:
407423
previous_sandbox_id = self.registry.get_current_sandbox_id(session_id)
@@ -411,7 +427,7 @@ def _set_current_sandbox_after_lease(
411427
self.registry.release_lease(previous_sandbox_id)
412428
self.registry.set_current_sandbox_id(session_id, sandbox_id)
413429
self.registry.touch_sandbox(sandbox_id)
414-
self.save_registry()
430+
await self.save_registry_async()
415431
return self.registry.get_sandbox(sandbox_id) or record
416432

417433
def get_current_sandbox(self, session_id: str) -> dict:
@@ -497,7 +513,7 @@ async def destroy_sandbox(self, session_id: str, sandbox_id: str) -> dict:
497513
self.clear_idle_state(sandbox_id)
498514
self.registry.delete_sandbox(sandbox_id)
499515
self.drop_boot_lock(sandbox_id)
500-
self.save_registry()
516+
await self.save_registry_async()
501517
return record
502518

503519
async def get_observer_booter_by_id(
@@ -516,10 +532,10 @@ async def get_observer_booter_by_id(
516532
if not await self.booter_available(booter):
517533
self.session_booter.pop(sandbox_id, None)
518534
self.registry.update_sandbox_status(sandbox_id, "unknown")
519-
self.save_registry()
535+
await self.save_registry_async()
520536
raise RuntimeError(f"Sandbox {sandbox_id} is not running")
521537
self.registry.touch_sandbox(sandbox_id)
522-
self.save_registry()
538+
await self.save_registry_async()
523539
idle_timeout = record.get("idle_timeout") or 0
524540
self.schedule_idle_cleanup(sandbox_id, float(idle_timeout))
525541
return booter
@@ -530,7 +546,7 @@ async def reconcile_on_startup(self) -> None:
530546
self.session_booter.clear()
531547
for sandbox_id in list(self.idle_state):
532548
self.clear_idle_state(sandbox_id)
533-
self.registry.save()
549+
await self.save_registry_async()
534550

535551
async def cleanup_managed_sandboxes(self) -> None:
536552
managed_records = self.list_sandboxes()
@@ -546,7 +562,7 @@ async def cleanup_managed_sandboxes(self) -> None:
546562
provider = self.get_provider(record.get("provider", ""))
547563
except RuntimeError as provider_error:
548564
self.registry.update_sandbox_status(sandbox_id, "unknown")
549-
self.save_registry()
565+
await self.save_registry_async()
550566
logger.warning(
551567
"[Computer] Skip managed sandbox cleanup for unsupported provider: sandbox_id=%s error=%s",
552568
sandbox_id,
@@ -560,7 +576,7 @@ async def cleanup_managed_sandboxes(self) -> None:
560576
self.session_booter.pop(sandbox_id, None)
561577
except Exception as shutdown_err:
562578
self.registry.update_sandbox_status(sandbox_id, "unknown")
563-
self.save_registry()
579+
await self.save_registry_async()
564580
logger.warning(
565581
"[Computer] Failed to shutdown managed sandbox %s: %s",
566582
sandbox_id,
@@ -570,7 +586,7 @@ async def cleanup_managed_sandboxes(self) -> None:
570586
self.clear_idle_state(sandbox_id)
571587
self.registry.delete_sandbox(sandbox_id)
572588
self.drop_boot_lock(sandbox_id)
573-
self.registry.save()
589+
await self.save_registry_async()
574590

575591
def clear_idle_state(self, sandbox_id: str) -> None:
576592
state = self.idle_state.pop(sandbox_id, None)
@@ -628,7 +644,7 @@ async def _expire_when_idle(
628644
except Exception as shutdown_err:
629645
self.session_booter[sandbox_id] = booter
630646
self.registry.update_sandbox_status(sandbox_id, "unknown")
631-
self.save_registry()
647+
await self.save_registry_async()
632648
logger.warning(
633649
"[Computer] Failed to shutdown idle sandbox %s: %s",
634650
sandbox_id,
@@ -640,7 +656,7 @@ async def _expire_when_idle(
640656
else:
641657
self.registry.delete_sandbox(sandbox_id)
642658
self.drop_boot_lock(sandbox_id)
643-
self.save_registry()
659+
await self.save_registry_async()
644660
return
645661
except asyncio.CancelledError:
646662
raise

astrbot/core/computer/sandbox_provider.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -10,6 +10,7 @@ class SandboxProvider(Protocol):
1010
provider_id: str
1111
capabilities: set[str]
1212
tool_names: set[str]
13+
system_prompt: str
1314

1415
def build_create_config(self, context: Context, session_id: str) -> dict: ...
1516

astrbot/core/computer/sandbox_registry.py

Lines changed: 10 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,6 @@
11
from __future__ import annotations
22

3+
import asyncio
34
import json
45
import time
56
from copy import deepcopy
@@ -327,9 +328,16 @@ def load(self) -> None:
327328
self._payload.get("session_current") or {}
328329
)
329330

330-
def save(self) -> None:
331+
def _write_payload(self, payload: dict[str, Any]) -> None:
331332
self.storage_path.parent.mkdir(parents=True, exist_ok=True)
332333
self.storage_path.write_text(
333-
json.dumps(self._payload, ensure_ascii=False, indent=2, sort_keys=True),
334+
json.dumps(payload, ensure_ascii=False, indent=2, sort_keys=True),
334335
encoding="utf-8",
335336
)
337+
338+
def save(self) -> None:
339+
self._write_payload(deepcopy(self._payload))
340+
341+
async def save_async(self) -> None:
342+
payload = deepcopy(self._payload)
343+
await asyncio.to_thread(self._write_payload, payload)
Lines changed: 23 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,23 @@
1+
from __future__ import annotations
2+
3+
from typing import Any
4+
5+
6+
def resolve_sandbox_provider_bindings(
7+
provider_id: str | None,
8+
tool_mgr: Any,
9+
) -> tuple[dict | None, list[Any]]:
10+
"""Return provider metadata and active provider tools for sandbox mode."""
11+
from astrbot.core.computer.computer_client import get_sandbox_provider_info
12+
13+
normalized_provider_id = "" if provider_id is None else str(provider_id).lower()
14+
provider_info = get_sandbox_provider_info(normalized_provider_id)
15+
if not provider_info:
16+
return None, []
17+
18+
tools = []
19+
for tool_name in provider_info.get("tool_names", []):
20+
tool = tool_mgr.get_func(tool_name)
21+
if tool and getattr(tool, "active", True):
22+
tools.append(tool)
23+
return provider_info, tools

0 commit comments

Comments
 (0)