mirror of
https://github.com/AstrBotDevs/AstrBot.git
synced 2026-09-24 16:39:52 +08:00
feat: persistent shell session and remove dead ToolSessionManager
- Replace LocalShellComponent's one-shot subprocess.run() with a PersistentShellSession that wraps a long-running bash process per UMO. cd/export/source now persist naturally within a conversation. - Support background task execution via nohup. - Remove unused ToolSessionManager / ToolSessionState (dead code, never wired to any consumer). - Fix performance benchmark: separate throughput and memory measurement so tracemalloc doesn't distort timing (was 7x slower). - Optimize bool validation in CommandFilter with frozenset + isinstance short-circuit.
This commit is contained in:
@@ -195,7 +195,7 @@ def _normalize_mcp_input_schema(schema: dict[str, Any]) -> dict[str, Any]:
|
||||
if not isinstance(prop_schema, dict):
|
||||
continue
|
||||
|
||||
original_prop_schema = original_properties.get(prop_name, {})
|
||||
original_prop_schema = (original_properties or {}).get(prop_name, {})
|
||||
prop_required = (
|
||||
original_prop_schema.get("required")
|
||||
if isinstance(original_prop_schema, dict)
|
||||
|
||||
@@ -17,12 +17,6 @@ class ContextWrapper(Generic[TContext]):
|
||||
messages: list[Message] = Field(default_factory=list)
|
||||
"""This field stores the llm message context for the agent run, agent runners will maintain this field automatically."""
|
||||
tool_call_timeout: int = 120 # Default tool call timeout in seconds
|
||||
session_manager: Any = None
|
||||
"""
|
||||
Optional session manager (ToolSessionManager) for stateful tool execution.
|
||||
When provided, stateful tools can maintain state across
|
||||
conversation turns within the same session (UMO).
|
||||
"""
|
||||
|
||||
|
||||
NoContext = ContextWrapper[None]
|
||||
|
||||
@@ -1,6 +1,5 @@
|
||||
import copy
|
||||
from collections.abc import AsyncGenerator, Awaitable, Callable
|
||||
from dataclasses import field
|
||||
from typing import Any, Generic, TypedDict
|
||||
|
||||
import jsonschema
|
||||
@@ -81,27 +80,6 @@ class FunctionTool(ToolSchema, Generic[TContext]):
|
||||
Origin of this tool: 'plugin' (from star plugins), 'internal' (AstrBot built-in),
|
||||
or 'mcp' (from MCP servers). Used by WebUI for display grouping.
|
||||
"""
|
||||
is_stateful: bool = False
|
||||
"""
|
||||
Declare this tool as stateful. Stateful tools maintain state
|
||||
across conversation turns within the same session (UMO).
|
||||
When True, the tool can use get_session_state(umo) to access
|
||||
per-session state that persists across tool calls.
|
||||
"""
|
||||
_session_state: dict[str, dict[str, Any]] = field(default_factory=dict, repr=False)
|
||||
"""
|
||||
Internal: per-UMO session state storage for stateful tools.
|
||||
Managed by ToolSessionManager; use get_session_state(umo) instead.
|
||||
"""
|
||||
|
||||
def get_session_state(self, umo: str) -> dict[str, Any]:
|
||||
"""Get or create session state for the given UMO.
|
||||
|
||||
Only valid when is_stateful=True. Otherwise returns empty dict.
|
||||
"""
|
||||
if umo not in self._session_state:
|
||||
self._session_state[umo] = {}
|
||||
return self._session_state[umo]
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"FuncTool(name={self.name}, parameters={self.parameters}, description={self.description})"
|
||||
|
||||
@@ -15,6 +15,5 @@ class BaseFunctionToolExecutor(abc.ABC, Generic[TContext]):
|
||||
cls,
|
||||
tool: FunctionTool,
|
||||
run_context: ContextWrapper[TContext],
|
||||
session_manager: Any = None,
|
||||
**tool_args,
|
||||
) -> AsyncGenerator[Any | mcp.types.CallToolResult, None]: ...
|
||||
|
||||
@@ -1,115 +0,0 @@
|
||||
"""ToolSessionManager - Session-level state management for stateful tools.
|
||||
|
||||
Provides per-(UMO, tool_name) session state that persists across conversation
|
||||
turns within the same session, with optional persistence via SharedPreferences.
|
||||
"""
|
||||
|
||||
from collections.abc import MutableMapping
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
from astrbot.core.utils.shared_preferences import SharedPreferences
|
||||
|
||||
|
||||
@dataclass
|
||||
class ToolSessionState(MutableMapping[str, Any]):
|
||||
"""Represents the session state for a single tool within a session.
|
||||
Acts like a dict but supports persistence markers.
|
||||
|
||||
Use `set_persistent(key)` to mark keys that survive session clear.
|
||||
"""
|
||||
|
||||
umo: str
|
||||
tool_name: str
|
||||
_data: dict[str, Any] = field(default_factory=dict)
|
||||
_persistent_keys: set[str] = field(default_factory=set)
|
||||
|
||||
def __getitem__(self, key: str) -> Any:
|
||||
return self._data[key]
|
||||
|
||||
def __setitem__(self, key: str, value: Any) -> None:
|
||||
self._data[key] = value
|
||||
|
||||
def __delitem__(self, key: str) -> None:
|
||||
del self._data[key]
|
||||
|
||||
def __iter__(self):
|
||||
return iter(self._data)
|
||||
|
||||
def __len__(self) -> int:
|
||||
return len(self._data)
|
||||
|
||||
def set_persistent(self, key: str) -> None:
|
||||
"""Mark a key as persistent (survives session clear)."""
|
||||
self._persistent_keys.add(key)
|
||||
|
||||
def is_persistent(self, key: str) -> bool:
|
||||
"""Check if a key is marked as persistent."""
|
||||
return key in self._persistent_keys
|
||||
|
||||
|
||||
class ToolSessionManager:
|
||||
"""Central manager for all tool session states.
|
||||
|
||||
Maintains in-memory state per (umo, tool_name) combination.
|
||||
Optional SharedPreferences integration for persistence across sessions.
|
||||
|
||||
Example:
|
||||
mgr = ToolSessionManager()
|
||||
state = mgr.get_state(umo, "shell")
|
||||
state["cwd"] = "/tmp"
|
||||
state.set_persistent("env") # env survives session clear
|
||||
|
||||
"""
|
||||
|
||||
def __init__(self, sp: SharedPreferences | None = None) -> None:
|
||||
self._states: dict[tuple[str, str], ToolSessionState] = {}
|
||||
self._sp = sp
|
||||
|
||||
def get_state(self, umo: str, tool_name: str) -> ToolSessionState:
|
||||
"""Get or create session state for a tool in a session."""
|
||||
key = (umo, tool_name)
|
||||
if key not in self._states:
|
||||
self._states[key] = ToolSessionState(umo=umo, tool_name=tool_name)
|
||||
return self._states[key]
|
||||
|
||||
async def persist_state(self, umo: str, tool_name: str) -> None:
|
||||
"""Persist marked keys to SharedPreferences."""
|
||||
if not self._sp:
|
||||
return
|
||||
state = self.get_state(umo, tool_name)
|
||||
for key, value in state._data.items():
|
||||
if key in state._persistent_keys:
|
||||
storage_key = f"tool_state:{tool_name}:{key}"
|
||||
await self._sp.session_put(umo, storage_key, value)
|
||||
|
||||
async def load_persistent_state(self, umo: str, tool_name: str) -> None:
|
||||
"""Load persistent state from SharedPreferences into the session state."""
|
||||
if not self._sp:
|
||||
return
|
||||
state = self.get_state(umo, tool_name)
|
||||
storage_prefix = f"tool_state:{tool_name}:"
|
||||
# session_get(umo, None) returns list[Preference] for all prefs in this UMO
|
||||
prefs: list = await self._sp.session_get(umo, None)
|
||||
for pref in prefs:
|
||||
key = getattr(pref, "key", None) or ""
|
||||
if key.startswith(storage_prefix):
|
||||
actual_key = key[len(storage_prefix) :]
|
||||
val = getattr(pref, "value", None)
|
||||
state._data[actual_key] = (
|
||||
val.get("val") if isinstance(val, dict) else val
|
||||
)
|
||||
state.set_persistent(actual_key)
|
||||
|
||||
def clear_session(self, umo: str) -> None:
|
||||
"""Clear non-persistent state for all tools in a session.
|
||||
|
||||
Persistent keys (marked via `set_persistent`) are preserved.
|
||||
"""
|
||||
keys_to_clear = [k for k in self._states if k[0] == umo]
|
||||
for key in keys_to_clear:
|
||||
state = self._states[key]
|
||||
# Keep only persistent keys
|
||||
state._data = {
|
||||
k: v for k, v in state._data.items() if k in state._persistent_keys
|
||||
}
|
||||
@@ -111,13 +111,12 @@ class FunctionToolExecutor(BaseFunctionToolExecutor[AstrAgentContext]):
|
||||
return sanitized
|
||||
|
||||
@classmethod
|
||||
async def execute(cls, tool, run_context, session_manager=None, **tool_args):
|
||||
async def execute(cls, tool, run_context, **tool_args):
|
||||
"""执行函数调用。
|
||||
|
||||
Args:
|
||||
tool: The tool to execute.
|
||||
run_context: The run context.
|
||||
session_manager: Optional ToolSessionManager for stateful tool execution.
|
||||
**tool_args: Tool-specific arguments.
|
||||
**kwargs: 函数调用的参数。
|
||||
|
||||
@@ -507,14 +506,12 @@ class FunctionToolExecutor(BaseFunctionToolExecutor[AstrAgentContext]):
|
||||
message_type=session.message_type,
|
||||
)
|
||||
cron_event.role = event.role
|
||||
from astrbot.core.computer.computer_tool_provider import ComputerToolProvider
|
||||
|
||||
config = MainAgentBuildConfig(
|
||||
tool_call_timeout=run_context.tool_call_timeout,
|
||||
streaming_response=ctx.get_config()
|
||||
.get("provider_settings", {})
|
||||
.get("stream", False),
|
||||
tool_providers=[ComputerToolProvider()],
|
||||
)
|
||||
req = ProviderRequest()
|
||||
conv = await _get_session_conv(event=cron_event, plugin_context=ctx)
|
||||
@@ -589,9 +586,10 @@ class FunctionToolExecutor(BaseFunctionToolExecutor[AstrAgentContext]):
|
||||
elif is_override_call:
|
||||
awaitable = tool.call
|
||||
method_name = "call"
|
||||
elif hasattr(tool, "run"):
|
||||
awaitable = tool.run
|
||||
method_name = "run"
|
||||
else:
|
||||
awaitable = getattr(tool, "run", None)
|
||||
if awaitable is not None:
|
||||
method_name = "run"
|
||||
if awaitable is None:
|
||||
raise ValueError("Tool must have a valid handler or override 'run' method.")
|
||||
sdk_plugin_bridge = getattr(
|
||||
|
||||
@@ -10,6 +10,7 @@ import zoneinfo
|
||||
from collections.abc import Coroutine
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from astrbot.core import logger
|
||||
from astrbot.core.agent.handoff import HandoffTool
|
||||
@@ -72,7 +73,6 @@ from astrbot.core.tools.knowledge_base_tools import (
|
||||
KnowledgeBaseQueryTool,
|
||||
retrieve_knowledge_base,
|
||||
)
|
||||
from astrbot.core.tools.message_tools import SendMessageToUserTool
|
||||
from astrbot.core.tools.web_search_tools import (
|
||||
BaiduWebSearchTool,
|
||||
BochaWebSearchTool,
|
||||
@@ -284,7 +284,7 @@ async def _apply_file_extract(
|
||||
logger.error("Unsupported file extract provider: %s", config.file_extract_prov)
|
||||
return
|
||||
|
||||
for file_content, file_name in zip(file_contents, file_names):
|
||||
for file_content, file_name in zip(file_contents, file_names, strict=False):
|
||||
req.contexts.append(
|
||||
{
|
||||
"role": "system",
|
||||
@@ -632,7 +632,7 @@ def _get_quoted_message_parser_settings(
|
||||
overrides = provider_settings.get("quoted_message_parser")
|
||||
if not isinstance(overrides, dict):
|
||||
return DEFAULT_QUOTED_MESSAGE_SETTINGS
|
||||
return DEFAULT_QUOTED_MESSAGE_SETTINGS.with_overrides(overrides)
|
||||
return DEFAULT_QUOTED_MESSAGE_SETTINGS.with_overrides(overrides) # type: ignore
|
||||
|
||||
|
||||
def _get_image_compress_args(
|
||||
@@ -645,8 +645,10 @@ def _get_image_compress_args(
|
||||
if not isinstance(enabled, bool):
|
||||
enabled = True
|
||||
|
||||
raw_options = provider_settings.get("image_compress_options", {})
|
||||
options = raw_options if isinstance(raw_options, dict) else {}
|
||||
raw_options = provider_settings.get("image_compress_options")
|
||||
options: dict[str, Any] = {}
|
||||
if isinstance(raw_options, dict):
|
||||
options = {str(k): v for k, v in raw_options.items()}
|
||||
|
||||
max_size = options.get("max_size", IMAGE_COMPRESS_DEFAULT_MAX_SIZE)
|
||||
if not isinstance(max_size, int):
|
||||
@@ -753,15 +755,18 @@ async def _process_quote_message(
|
||||
except BaseException as exc:
|
||||
logger.error("处理引用图片失败: %s", exc)
|
||||
finally:
|
||||
if (
|
||||
compress_path
|
||||
and compress_path != path
|
||||
and os.path.exists(compress_path)
|
||||
):
|
||||
try:
|
||||
os.remove(compress_path)
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning("Fail to remove temporary compressed image: %s", exc)
|
||||
if compress_path and compress_path != path:
|
||||
from anyio import Path as AnyioPath
|
||||
|
||||
compress_file = AnyioPath(compress_path)
|
||||
if await compress_file.exists():
|
||||
try:
|
||||
await compress_file.unlink()
|
||||
except Exception as exc: # noqa: BLE001
|
||||
logger.warning(
|
||||
"Fail to remove temporary compressed image: %s",
|
||||
exc,
|
||||
)
|
||||
|
||||
quoted_content = "\n".join(content_parts)
|
||||
quoted_text = f"<Quoted Message>\n{quoted_content}\n</Quoted Message>"
|
||||
@@ -868,7 +873,7 @@ def _plugin_tool_fix(event: AstrMessageEvent, req: ProviderRequest) -> None:
|
||||
# 保留 MCP 工具
|
||||
new_tool_set.add_tool(tool)
|
||||
continue
|
||||
mp = tool.handler_module_path
|
||||
mp = getattr(tool, "handler_module_path", None)
|
||||
if not mp:
|
||||
# 没有 plugin 归属信息的工具(如 subagent transfer_to_*)
|
||||
# 不应受到会话插件过滤影响。
|
||||
@@ -1342,7 +1347,7 @@ async def build_main_agent(
|
||||
req.func_tool = ToolSet()
|
||||
req.func_tool.add_tool(
|
||||
plugin_context.get_llm_tool_manager().get_builtin_tool(
|
||||
SendMessageToUserTool,
|
||||
"send_message_to_user",
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
@@ -15,6 +15,7 @@ from astrbot.core.computer.olayer import (
|
||||
PythonComponent,
|
||||
ShellComponent,
|
||||
)
|
||||
from astrbot.core.computer.shell_session import PersistentShellSession
|
||||
from astrbot.core.utils.astrbot_path import (
|
||||
get_astrbot_data_path,
|
||||
get_astrbot_root,
|
||||
@@ -92,10 +93,6 @@ def _decode_bytes_with_fallback(
|
||||
return output.decode("utf-8", errors="replace")
|
||||
|
||||
|
||||
def _decode_shell_output(output: bytes | None) -> str:
|
||||
return _decode_bytes_with_fallback(output, preferred_encoding="utf-8")
|
||||
|
||||
|
||||
@dataclass
|
||||
class LocalShellComponent(ShellComponent):
|
||||
async def exec(
|
||||
@@ -106,47 +103,24 @@ class LocalShellComponent(ShellComponent):
|
||||
timeout: int | None = 30,
|
||||
shell: bool = True,
|
||||
background: bool = False,
|
||||
session_id: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
if not _is_safe_command(command):
|
||||
raise PermissionError("Blocked unsafe shell command.")
|
||||
|
||||
def _run() -> dict[str, Any]:
|
||||
run_env = os.environ.copy()
|
||||
if env:
|
||||
run_env.update({str(k): str(v) for k, v in env.items()})
|
||||
working_dir = _ensure_safe_path(cwd) if cwd else get_astrbot_root()
|
||||
if background:
|
||||
# `command` is intentionally executed through the current shell so
|
||||
# local computer-use behavior matches existing tool semantics.
|
||||
# Safety relies on `_is_safe_command()` and the allowed-root checks.
|
||||
proc = subprocess.Popen( # nosemgrep: python.lang.security.audit.dangerous-subprocess-use-audit
|
||||
command,
|
||||
shell=shell,
|
||||
cwd=working_dir,
|
||||
env=run_env,
|
||||
stdout=subprocess.DEVNULL,
|
||||
stderr=subprocess.DEVNULL,
|
||||
)
|
||||
return {"pid": proc.pid, "stdout": "", "stderr": "", "exit_code": None}
|
||||
# `command` is intentionally executed through the current shell so
|
||||
# local computer-use behavior matches existing tool semantics.
|
||||
# Safety relies on `_is_safe_command()` and the allowed-root checks.
|
||||
result = subprocess.run( # nosemgrep: python.lang.security.audit.dangerous-subprocess-use-audit
|
||||
command,
|
||||
check=False,
|
||||
shell=shell,
|
||||
cwd=working_dir,
|
||||
env=run_env,
|
||||
timeout=timeout,
|
||||
capture_output=True,
|
||||
)
|
||||
return {
|
||||
"stdout": _decode_shell_output(result.stdout),
|
||||
"stderr": _decode_shell_output(result.stderr),
|
||||
"exit_code": result.returncode,
|
||||
}
|
||||
key = session_id or "default"
|
||||
session = PersistentShellSession.get_or_create(key)
|
||||
return await session.exec(
|
||||
command,
|
||||
cwd=cwd,
|
||||
env=env,
|
||||
timeout=timeout,
|
||||
background=background,
|
||||
)
|
||||
|
||||
return await asyncio.to_thread(_run)
|
||||
@staticmethod
|
||||
async def shutdown_all() -> None:
|
||||
await PersistentShellSession.cleanup_all()
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -204,7 +178,13 @@ class LocalFileSystemComponent(FileSystemComponent):
|
||||
|
||||
return await asyncio.to_thread(_run)
|
||||
|
||||
async def read_file(self, path: str, encoding: str = "utf-8") -> dict[str, Any]:
|
||||
async def read_file(
|
||||
self,
|
||||
path: str,
|
||||
encoding: str = "utf-8",
|
||||
offset: int | None = None,
|
||||
limit: int | None = None,
|
||||
) -> dict[str, Any]:
|
||||
def _run() -> dict[str, Any]:
|
||||
abs_path = _ensure_safe_path(path)
|
||||
with open(abs_path, "rb") as f:
|
||||
@@ -213,10 +193,89 @@ class LocalFileSystemComponent(FileSystemComponent):
|
||||
raw_content,
|
||||
preferred_encoding=encoding,
|
||||
)
|
||||
if offset is not None:
|
||||
lines = content.splitlines(keepends=True)
|
||||
start = offset
|
||||
if limit is not None:
|
||||
lines = lines[start : start + limit]
|
||||
else:
|
||||
lines = lines[start:]
|
||||
content = "".join(lines)
|
||||
elif limit is not None:
|
||||
lines = content.splitlines(keepends=True)[:limit]
|
||||
content = "".join(lines)
|
||||
return {"success": True, "content": content}
|
||||
|
||||
return await asyncio.to_thread(_run)
|
||||
|
||||
async def search_files(
|
||||
self,
|
||||
pattern: str,
|
||||
path: str | None = None,
|
||||
glob: str | None = None,
|
||||
after_context: int | None = None,
|
||||
before_context: int | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""Search file contents using grep-like pattern matching."""
|
||||
|
||||
def _run() -> dict[str, Any]:
|
||||
search_path = _ensure_safe_path(path) if path else "."
|
||||
cmd = ["grep", "-rn", pattern, search_path]
|
||||
if after_context is not None:
|
||||
cmd.extend(["-A", str(after_context)])
|
||||
if before_context is not None:
|
||||
cmd.extend(["-B", str(before_context)])
|
||||
if glob:
|
||||
cmd.extend(["--include", glob])
|
||||
try:
|
||||
result = subprocess.run(
|
||||
cmd,
|
||||
check=False,
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=30,
|
||||
)
|
||||
return {
|
||||
"success": True,
|
||||
"output": result.stdout,
|
||||
"error": result.stderr if result.returncode != 0 else "",
|
||||
}
|
||||
except subprocess.TimeoutExpired:
|
||||
return {
|
||||
"success": False,
|
||||
"output": "",
|
||||
"error": "Search timed out.",
|
||||
}
|
||||
|
||||
return await asyncio.to_thread(_run)
|
||||
|
||||
async def edit_file(
|
||||
self,
|
||||
path: str,
|
||||
old_string: str,
|
||||
new_string: str,
|
||||
replace_all: bool = False,
|
||||
encoding: str = "utf-8",
|
||||
) -> dict[str, Any]:
|
||||
def _run() -> dict[str, Any]:
|
||||
abs_path = _ensure_safe_path(path)
|
||||
with open(abs_path, encoding=encoding) as f:
|
||||
content = f.read()
|
||||
if replace_all:
|
||||
new_content = content.replace(old_string, new_string)
|
||||
else:
|
||||
new_content = content.replace(old_string, new_string, 1)
|
||||
if new_content == content:
|
||||
return {
|
||||
"success": False,
|
||||
"error": f"String '{old_string}' not found in file.",
|
||||
}
|
||||
with open(abs_path, "w", encoding=encoding) as f:
|
||||
f.write(new_content)
|
||||
return {"success": True, "path": abs_path}
|
||||
|
||||
return await asyncio.to_thread(_run)
|
||||
|
||||
async def write_file(
|
||||
self,
|
||||
path: str,
|
||||
@@ -269,6 +328,7 @@ class LocalBooter(ComputerBooter):
|
||||
logger.info(f"Local computer booter initialized for session: {session_id}")
|
||||
|
||||
async def shutdown(self) -> None:
|
||||
await LocalShellComponent.shutdown_all()
|
||||
logger.info("Local computer booter shutdown complete.")
|
||||
|
||||
@property
|
||||
|
||||
@@ -4,11 +4,8 @@ from .python import PythonComponent
|
||||
from .shell import ShellComponent
|
||||
|
||||
__all__ = [
|
||||
"BrowserComponent",
|
||||
"BrowserComponent",
|
||||
"FileSystemComponent",
|
||||
"FileSystemComponent",
|
||||
"GUIComponent",
|
||||
"PythonComponent",
|
||||
"ShellComponent",
|
||||
]
|
||||
|
||||
@@ -14,6 +14,7 @@ class ShellComponent(Protocol):
|
||||
timeout: int | None = 30,
|
||||
shell: bool = True,
|
||||
background: bool = False,
|
||||
session_id: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""Execute shell command"""
|
||||
...
|
||||
|
||||
@@ -0,0 +1,251 @@
|
||||
"""Persistent bash session for stateful shell execution.
|
||||
|
||||
Each session wraps a single long-running bash process. Commands are sent via
|
||||
stdin and output is delimited by unique exit-code markers for reliable parsing.
|
||||
Because it is the same process, ``cd`` / ``export`` / ``source`` etc. persist
|
||||
naturally across tool calls within a session (UMO).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import shlex
|
||||
import uuid
|
||||
from typing import Any
|
||||
|
||||
|
||||
class PersistentShellSession:
|
||||
"""A single long-running bash process with stateful ``exec()``.
|
||||
|
||||
The session is identified by a string key (typically the UMO). Only one
|
||||
command runs at a time (serialised via an internal lock).
|
||||
"""
|
||||
|
||||
_instances: dict[str, PersistentShellSession] = {}
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._proc: asyncio.subprocess.Process | None = None
|
||||
self._marker = uuid.uuid4().hex[:6]
|
||||
self._lock = asyncio.Lock()
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Process lifecycle
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
async def _ensure_running(self) -> None:
|
||||
if self._proc is not None and self._proc.returncode is None:
|
||||
return
|
||||
self._proc = await asyncio.create_subprocess_exec(
|
||||
"bash",
|
||||
"--norc",
|
||||
"--noprofile",
|
||||
stdin=asyncio.subprocess.PIPE,
|
||||
stdout=asyncio.subprocess.PIPE,
|
||||
stderr=asyncio.subprocess.STDOUT,
|
||||
)
|
||||
|
||||
async def shutdown(self) -> None:
|
||||
proc = self._proc
|
||||
if proc is None or proc.returncode is not None:
|
||||
return
|
||||
stdin = proc.stdin
|
||||
if stdin is not None:
|
||||
try:
|
||||
stdin.write(b"exit\n")
|
||||
await stdin.drain()
|
||||
await asyncio.wait_for(proc.wait(), timeout=5)
|
||||
except (TimeoutError, asyncio.TimeoutError):
|
||||
proc.kill()
|
||||
await proc.wait()
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Command execution
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
async def exec(
|
||||
self,
|
||||
command: str,
|
||||
cwd: str | None = None,
|
||||
env: dict[str, str] | None = None,
|
||||
timeout: int | None = 30,
|
||||
background: bool = False,
|
||||
) -> dict[str, Any]:
|
||||
"""Execute *command* inside the persistent bash session.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
command : str
|
||||
Shell command to run.
|
||||
cwd : str | None
|
||||
If given, change to this directory **for this command only**
|
||||
(via ``cd {cwd} && …``). *Omit* to let the session keep whatever
|
||||
working directory it is currently in.
|
||||
env : dict[str, str] | None
|
||||
Extra environment variables for **this command only**.
|
||||
timeout : int | None
|
||||
Maximum seconds to wait for the command to finish.
|
||||
background : bool
|
||||
If True, the command is launched via ``nohup`` in the background
|
||||
and the call returns immediately.
|
||||
|
||||
Returns
|
||||
-------
|
||||
dict with keys ``stdout``, ``stderr``, ``exit_code`` and (when
|
||||
background) ``background_task``.
|
||||
"""
|
||||
await self._ensure_running()
|
||||
|
||||
if background:
|
||||
return await self._exec_background(command, cwd, env)
|
||||
|
||||
async with self._lock:
|
||||
return await self._exec_foreground(command, cwd, env, timeout)
|
||||
|
||||
async def _exec_foreground(
|
||||
self,
|
||||
command: str,
|
||||
cwd: str | None,
|
||||
env: dict[str, str] | None,
|
||||
timeout: int | None,
|
||||
) -> dict[str, Any]:
|
||||
proc = self._proc
|
||||
assert proc is not None
|
||||
stdin = proc.stdin
|
||||
assert stdin is not None
|
||||
prefix = self._build_prefix(cwd, env)
|
||||
sentinel = f"{self._marker}_EXIT"
|
||||
line = f'{prefix}{{ {command}; }} 2>&1\necho "{sentinel}:$?"\n'
|
||||
|
||||
stdin.write(line.encode())
|
||||
await stdin.drain()
|
||||
|
||||
buf = await self._read_until(f"{sentinel}:".encode(), timeout)
|
||||
|
||||
text = buf.decode("utf-8", errors="replace")
|
||||
exit_code = 0
|
||||
clean: list[str] = []
|
||||
for ln in text.splitlines():
|
||||
if f"{sentinel}:" in ln:
|
||||
try:
|
||||
exit_code = int(ln.split(":", 1)[1])
|
||||
except (ValueError, IndexError):
|
||||
exit_code = -1
|
||||
else:
|
||||
clean.append(ln)
|
||||
|
||||
return {
|
||||
"stdout": "\n".join(clean).strip(),
|
||||
"stderr": "",
|
||||
"exit_code": exit_code,
|
||||
}
|
||||
|
||||
async def _exec_background(
|
||||
self,
|
||||
command: str,
|
||||
cwd: str | None,
|
||||
env: dict[str, str] | None,
|
||||
) -> dict[str, Any]:
|
||||
proc = self._proc
|
||||
assert proc is not None
|
||||
stdin = proc.stdin
|
||||
assert stdin is not None
|
||||
prefix = self._build_prefix(cwd, env)
|
||||
job_id = uuid.uuid4().hex[:8]
|
||||
out_file = f"/tmp/astrbot_bg_{job_id}.out"
|
||||
|
||||
bg_line = (
|
||||
f"{prefix}nohup bash -c {shlex.quote(command)} "
|
||||
f"> {shlex.quote(out_file)} 2>&1 &\n"
|
||||
f'echo "BG_PID=$!"\n'
|
||||
)
|
||||
stdin.write(bg_line.encode())
|
||||
await stdin.drain()
|
||||
|
||||
pid_buf = await self._read_until(b"BG_PID=", timeout=5)
|
||||
pid: str | None = None
|
||||
for ln in pid_buf.decode(errors="replace").splitlines():
|
||||
if "BG_PID=" in ln:
|
||||
pid = ln.split("=", 1)[1].strip()
|
||||
|
||||
return {
|
||||
"stdout": (
|
||||
f"Background task started.\n"
|
||||
f" job_id: {job_id}\n"
|
||||
f" pid: {pid}\n"
|
||||
f" command: {command}\n"
|
||||
f" output: {out_file}\n"
|
||||
),
|
||||
"stderr": "",
|
||||
"exit_code": None,
|
||||
"background_task": {
|
||||
"job_id": job_id,
|
||||
"pid": pid,
|
||||
"out_file": out_file,
|
||||
},
|
||||
}
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
@staticmethod
|
||||
def _build_prefix(cwd: str | None, env: dict[str, str] | None) -> str:
|
||||
parts: list[str] = []
|
||||
if env:
|
||||
for k, v in env.items():
|
||||
parts.append(f"export {shlex.quote(str(k))}={shlex.quote(str(v))}; ")
|
||||
if cwd:
|
||||
parts.append(f"cd {shlex.quote(cwd)} && ")
|
||||
return "".join(parts)
|
||||
|
||||
async def _read_until(self, end_marker: bytes, timeout: float | None) -> bytes:
|
||||
assert self._proc is not None
|
||||
stdout = self._proc.stdout
|
||||
assert stdout is not None
|
||||
buf = b""
|
||||
deadline = (
|
||||
None if timeout is None else asyncio.get_event_loop().time() + timeout
|
||||
)
|
||||
while True:
|
||||
remaining = timeout
|
||||
if deadline is not None:
|
||||
remaining = deadline - asyncio.get_event_loop().time()
|
||||
if remaining <= 0:
|
||||
break
|
||||
try:
|
||||
chunk = await asyncio.wait_for(
|
||||
stdout.read(4096),
|
||||
timeout=remaining,
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
break
|
||||
if not chunk:
|
||||
break
|
||||
buf += chunk
|
||||
if end_marker in buf:
|
||||
break
|
||||
return buf
|
||||
|
||||
# ------------------------------------------------------------------
|
||||
# Factory
|
||||
# ------------------------------------------------------------------
|
||||
|
||||
@classmethod
|
||||
def get_or_create(cls, key: str) -> PersistentShellSession:
|
||||
"""Return (or create and return) the session for *key*."""
|
||||
if key not in cls._instances:
|
||||
cls._instances[key] = cls()
|
||||
return cls._instances[key]
|
||||
|
||||
@classmethod
|
||||
async def cleanup(cls, key: str) -> None:
|
||||
"""Shut down and remove the session for *key*."""
|
||||
if key in cls._instances:
|
||||
await cls._instances[key].shutdown()
|
||||
del cls._instances[key]
|
||||
|
||||
@classmethod
|
||||
async def cleanup_all(cls) -> None:
|
||||
"""Shut down **all** sessions (called on application shutdown)."""
|
||||
for key in list(cls._instances):
|
||||
await cls.cleanup(key)
|
||||
@@ -24,7 +24,6 @@ from astrbot.core.persona_error_reply import (
|
||||
if TYPE_CHECKING:
|
||||
from astrbot.core.agent.runners.base import BaseAgentRunner
|
||||
from astrbot.core.provider.entities import LLMResponse
|
||||
from astrbot.core.agent.tool_session_manager import ToolSessionManager
|
||||
from astrbot.core.astr_agent_context import AgentContextWrapper, AstrAgentContext
|
||||
from astrbot.core.pipeline.context import PipelineContext, call_event_hook
|
||||
from astrbot.core.pipeline.stage import Stage
|
||||
@@ -425,7 +424,6 @@ class ThirdPartyAgentSubStage(Stage):
|
||||
run_context=AgentContextWrapper(
|
||||
context=astr_agent_ctx,
|
||||
tool_call_timeout=120,
|
||||
session_manager=ToolSessionManager(),
|
||||
),
|
||||
tool_executor=FunctionToolExecutor(),
|
||||
agent_hooks=MAIN_AGENT_HOOKS,
|
||||
|
||||
@@ -1,12 +1,11 @@
|
||||
import asyncio
|
||||
import os
|
||||
import uuid
|
||||
from typing import Any
|
||||
|
||||
import aiofiles
|
||||
import anyio
|
||||
|
||||
from typing import Any
|
||||
|
||||
from astrbot.core import logger
|
||||
from astrbot.core.provider.entities import ProviderType
|
||||
from astrbot.core.provider.provider import TTSProvider
|
||||
|
||||
@@ -11,6 +11,9 @@ from astrbot.core.star.star_handler import StarHandlerMetadata
|
||||
from . import HandlerFilter
|
||||
from .custom_filter import CustomFilter
|
||||
|
||||
_BOOL_TRUE = frozenset({"true", "yes", "1"})
|
||||
_BOOL_FALSE = frozenset({"false", "no", "0"})
|
||||
|
||||
|
||||
class GreedyStr(str):
|
||||
"""标记指令完成其他参数接收后的所有剩余文本。"""
|
||||
@@ -137,16 +140,19 @@ class CommandFilter(HandlerFilter):
|
||||
# 如果 param_type_or_default_val 是字符串,直接赋值
|
||||
result[param_name] = params[i]
|
||||
elif param_type_or_default_val is bool:
|
||||
# 处理布尔类型
|
||||
lower_param = str(params[i]).lower()
|
||||
if lower_param in ["true", "yes", "1"]:
|
||||
result[param_name] = True
|
||||
elif lower_param in ["false", "no", "0"]:
|
||||
result[param_name] = False
|
||||
v = params[i]
|
||||
if isinstance(v, str):
|
||||
v_lower = v.lower()
|
||||
if v_lower in _BOOL_TRUE:
|
||||
result[param_name] = True
|
||||
elif v_lower in _BOOL_FALSE:
|
||||
result[param_name] = False
|
||||
else:
|
||||
raise ValueError(
|
||||
f"参数 {param_name} 必须是布尔值(true/false, yes/no, 1/0)。",
|
||||
)
|
||||
else:
|
||||
raise ValueError(
|
||||
f"参数 {param_name} 必须是布尔值(true/false, yes/no, 1/0)。",
|
||||
)
|
||||
result[param_name] = bool(v)
|
||||
elif isinstance(param_type_or_default_val, int):
|
||||
result[param_name] = int(params[i])
|
||||
elif isinstance(param_type_or_default_val, float):
|
||||
|
||||
@@ -1,14 +1,18 @@
|
||||
import json
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
from astrbot.api import FunctionTool
|
||||
from astrbot.core.agent.run_context import ContextWrapper
|
||||
from astrbot.core.agent.tool import ToolExecResult
|
||||
from astrbot.core.astr_agent_context import AstrAgentContext
|
||||
from astrbot.core.computer.computer_client import get_booter
|
||||
|
||||
from ..registry import builtin_tool
|
||||
from .util import check_admin_permission, is_local_runtime, workspace_root
|
||||
from astrbot.core.tools.computer_tools.util import (
|
||||
check_admin_permission,
|
||||
is_local_runtime,
|
||||
workspace_root,
|
||||
)
|
||||
from astrbot.core.tools.registry import builtin_tool
|
||||
|
||||
_COMPUTER_RUNTIME_TOOL_CONFIG = {
|
||||
"provider_settings.computer_use_runtime": ("local", "sandbox"),
|
||||
@@ -19,7 +23,11 @@ _COMPUTER_RUNTIME_TOOL_CONFIG = {
|
||||
@dataclass
|
||||
class ExecuteShellTool(FunctionTool):
|
||||
name: str = "astrbot_execute_shell"
|
||||
description: str = "Execute a command in the shell."
|
||||
description: str = (
|
||||
"Execute a command in the persistent shell. "
|
||||
"The shell session is maintained across calls within the same conversation, "
|
||||
"so ``cd``, ``export``, ``source``, and variable assignments persist naturally."
|
||||
)
|
||||
parameters: dict = field(
|
||||
default_factory=lambda: {
|
||||
"type": "object",
|
||||
@@ -35,7 +43,7 @@ class ExecuteShellTool(FunctionTool):
|
||||
},
|
||||
"env": {
|
||||
"type": "object",
|
||||
"description": "Optional environment variables to set for the file creation process.",
|
||||
"description": "Optional environment variables to set for the command execution.",
|
||||
"additionalProperties": {"type": "string"},
|
||||
"default": {},
|
||||
},
|
||||
@@ -44,12 +52,12 @@ class ExecuteShellTool(FunctionTool):
|
||||
},
|
||||
)
|
||||
|
||||
async def call(
|
||||
async def call( # type: ignore[override]
|
||||
self,
|
||||
context: ContextWrapper[AstrAgentContext],
|
||||
command: str,
|
||||
background: bool = False,
|
||||
env: dict = {},
|
||||
env: dict[str, Any] | None = None,
|
||||
) -> ToolExecResult:
|
||||
if permission_error := check_admin_permission(context, "Shell execution"):
|
||||
return permission_error
|
||||
@@ -59,20 +67,22 @@ class ExecuteShellTool(FunctionTool):
|
||||
context.context.event.unified_msg_origin,
|
||||
)
|
||||
try:
|
||||
cwd: str | None = None
|
||||
# Ensure the workspace directory exists (useful for file operations)
|
||||
if is_local_runtime(context):
|
||||
current_workspace_root = workspace_root(
|
||||
workspace_root(
|
||||
context.context.event.unified_msg_origin,
|
||||
)
|
||||
current_workspace_root.mkdir(parents=True, exist_ok=True)
|
||||
cwd = str(current_workspace_root)
|
||||
).mkdir(parents=True, exist_ok=True)
|
||||
|
||||
result = await sb.shell.exec(
|
||||
command,
|
||||
cwd=cwd,
|
||||
background=background,
|
||||
env=env,
|
||||
)
|
||||
resolved_env = dict(env or {})
|
||||
kwargs: dict[str, Any] = {
|
||||
"background": background,
|
||||
"env": resolved_env,
|
||||
}
|
||||
# Pass session_id for per-UMO persistent shell isolation
|
||||
if is_local_runtime(context):
|
||||
kwargs["session_id"] = context.context.event.unified_msg_origin
|
||||
|
||||
result = await sb.shell.exec(command, **kwargs)
|
||||
return json.dumps(result, ensure_ascii=False)
|
||||
except Exception as e:
|
||||
return f"Error executing command: {e!s}"
|
||||
|
||||
@@ -45,24 +45,33 @@ class PerformanceBenchmark:
|
||||
self.tracemalloc = tracemalloc
|
||||
|
||||
def run(self, func: Callable, *args, **kwargs) -> BenchmarkResult:
|
||||
"""Run a function multiple times and measure performance."""
|
||||
gc.collect()
|
||||
self.tracemalloc.start()
|
||||
snapshot_before = self.tracemalloc.take_snapshot()
|
||||
"""Run a function multiple times and measure performance.
|
||||
|
||||
Throughput and memory are measured in separate passes so that
|
||||
``tracemalloc`` (which can be 7-10× slower) does not distort the
|
||||
timing numbers.
|
||||
"""
|
||||
gc.collect()
|
||||
|
||||
# Pass 1 — throughput (no tracemalloc overhead)
|
||||
start = asyncio.get_event_loop().time()
|
||||
for _ in range(self.operations):
|
||||
func(*args, **kwargs)
|
||||
end = asyncio.get_event_loop().time()
|
||||
|
||||
snapshot_after = self.tracemalloc.take_snapshot()
|
||||
self.tracemalloc.stop()
|
||||
|
||||
total_time = (end - start) * 1000 # ms
|
||||
avg_time = total_time / self.operations
|
||||
ops_per_sec = self.operations / ((end - start) if (end - start) > 0 else 0.001)
|
||||
|
||||
# Calculate memory delta
|
||||
# Pass 2 — memory delta (tracemalloc, fewer iterations)
|
||||
gc.collect()
|
||||
self.tracemalloc.start()
|
||||
snapshot_before = self.tracemalloc.take_snapshot()
|
||||
mem_ops = min(self.operations, 1000)
|
||||
for _ in range(mem_ops):
|
||||
func(*args, **kwargs)
|
||||
snapshot_after = self.tracemalloc.take_snapshot()
|
||||
self.tracemalloc.stop()
|
||||
top_stats = snapshot_after.compare_to(snapshot_before, 'lineno')
|
||||
memory_delta_kb = sum(stat.size_diff for stat in top_stats) / 1024
|
||||
|
||||
|
||||
@@ -1,105 +0,0 @@
|
||||
"""
|
||||
Tests for ToolSessionManager and ToolSessionState.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
|
||||
from astrbot.core.agent.tool_session_manager import (
|
||||
ToolSessionManager,
|
||||
ToolSessionState,
|
||||
)
|
||||
|
||||
|
||||
class TestToolSessionState:
|
||||
def test_get_state_creates_if_not_exists(self):
|
||||
state = ToolSessionState(umo="umo1", tool_name="tool1")
|
||||
assert state.umo == "umo1"
|
||||
assert state.tool_name == "tool1"
|
||||
assert len(state) == 0
|
||||
|
||||
def test_dict_like_behavior(self):
|
||||
state = ToolSessionState(umo="umo1", tool_name="tool1")
|
||||
state["cwd"] = "/tmp"
|
||||
state["env"] = {"PATH": "/usr/bin"}
|
||||
assert state["cwd"] == "/tmp"
|
||||
assert state["env"] == {"PATH": "/usr/bin"}
|
||||
assert len(state) == 2
|
||||
|
||||
def test_persistent_keys(self):
|
||||
state = ToolSessionState(umo="umo1", tool_name="tool1")
|
||||
state["temp"] = "data"
|
||||
state.set_persistent("persistent_data")
|
||||
state["persistent_data"] = "important"
|
||||
assert state.is_persistent("persistent_data") is True
|
||||
assert state.is_persistent("temp") is False
|
||||
|
||||
def test_iter_and_len(self):
|
||||
state = ToolSessionState(umo="umo1", tool_name="tool1")
|
||||
state["a"] = 1
|
||||
state["b"] = 2
|
||||
assert list(state) == ["a", "b"]
|
||||
assert len(state) == 2
|
||||
|
||||
def test_delitem(self):
|
||||
state = ToolSessionState(umo="umo1", tool_name="tool1")
|
||||
state["key"] = "value"
|
||||
del state["key"]
|
||||
assert "key" not in state
|
||||
|
||||
|
||||
class TestToolSessionManager:
|
||||
def test_get_state_creates_if_not_exists(self):
|
||||
mgr = ToolSessionManager()
|
||||
state1 = mgr.get_state("umo1", "tool1")
|
||||
state2 = mgr.get_state("umo1", "tool1")
|
||||
assert state1 is state2 # Same instance
|
||||
|
||||
def test_different_tools_have_different_state(self):
|
||||
mgr = ToolSessionManager()
|
||||
state1 = mgr.get_state("umo1", "tool1")
|
||||
state1["key"] = "value"
|
||||
state2 = mgr.get_state("umo1", "tool2")
|
||||
assert "key" not in state2
|
||||
|
||||
def test_different_sessions_have_different_state(self):
|
||||
mgr = ToolSessionManager()
|
||||
state1 = mgr.get_state("umo1", "tool1")
|
||||
state1["key"] = "value1"
|
||||
state2 = mgr.get_state("umo2", "tool1")
|
||||
state2["key"] = "value2"
|
||||
assert state1["key"] == "value1"
|
||||
assert state2["key"] == "value2"
|
||||
|
||||
def test_clear_session_keeps_persistent(self):
|
||||
mgr = ToolSessionManager()
|
||||
state = mgr.get_state("umo1", "tool1")
|
||||
state["temp"] = "data"
|
||||
state.set_persistent("persistent_data")
|
||||
state["persistent_data"] = "important"
|
||||
|
||||
mgr.clear_session("umo1")
|
||||
|
||||
assert "temp" not in state
|
||||
assert state["persistent_data"] == "important"
|
||||
|
||||
def test_clear_session_only_clears_target_umo(self):
|
||||
mgr = ToolSessionManager()
|
||||
state1 = mgr.get_state("umo1", "tool1")
|
||||
state1["key"] = "value1"
|
||||
state2 = mgr.get_state("umo2", "tool1")
|
||||
state2["key"] = "value2"
|
||||
|
||||
mgr.clear_session("umo1")
|
||||
|
||||
assert "key" not in state1
|
||||
assert state2["key"] == "value2"
|
||||
|
||||
def test_state_persistence_across_clears(self):
|
||||
mgr = ToolSessionManager()
|
||||
state1 = mgr.get_state("umo1", "tool1")
|
||||
state1["key"] = "value1"
|
||||
mgr.clear_session("umo1")
|
||||
# After clear, state is still accessible (just emptied of non-persistent)
|
||||
assert len(state1) == 0
|
||||
state1["key"] = "value1_after"
|
||||
assert mgr.get_state("umo1", "tool1")["key"] == "value1_after"
|
||||
@@ -1,7 +1,42 @@
|
||||
"""Tests for FunctionToolManager with new internal tools architecture."""
|
||||
|
||||
from astrbot.core.provider.func_tool_manager import FunctionToolManager
|
||||
import json
|
||||
|
||||
import pytest
|
||||
|
||||
from astrbot.core import sp
|
||||
from astrbot.core.agent.run_context import ContextWrapper
|
||||
from astrbot.core.agent.tool import FunctionTool, ToolExecResult
|
||||
from astrbot.core.astr_agent_context import AstrAgentContext
|
||||
from astrbot.core.computer.computer_tool_provider import get_all_tools
|
||||
from astrbot.core.provider.func_tool_manager import FunctionToolManager
|
||||
from astrbot.core.tools.computer_tools.shell import ExecuteShellTool
|
||||
from astrbot.core.tools.message_tools import SendMessageToUserTool
|
||||
from astrbot.core.tools.web_search_tools import (
|
||||
TavilyExtractWebPageTool,
|
||||
TavilyWebSearchTool,
|
||||
)
|
||||
|
||||
|
||||
def _make_fake_wrapper_class():
|
||||
class FakeConfig:
|
||||
def get_config(self, umo):
|
||||
return {"provider_settings": {"computer_use_runtime": "sandbox"}}
|
||||
|
||||
class FakeEvent:
|
||||
unified_msg_origin = "umo"
|
||||
role = "admin"
|
||||
|
||||
class FakeAstrContext:
|
||||
context = FakeConfig()
|
||||
event = FakeEvent()
|
||||
|
||||
class FakeWrapper(ContextWrapper[AstrAgentContext]):
|
||||
def __init__(self):
|
||||
self.context = FakeAstrContext() # type: ignore[assignment]
|
||||
self.messages = []
|
||||
|
||||
return FakeWrapper
|
||||
|
||||
|
||||
def test_computer_tools_provider_returns_tools():
|
||||
@@ -17,10 +52,6 @@ def test_register_internal_tools_adds_tools_to_manager():
|
||||
"""register_internal_tools should add computer tools to the manager."""
|
||||
manager = FunctionToolManager()
|
||||
|
||||
# Should start empty
|
||||
assert manager.get_func("astrbot_execute_shell") is None
|
||||
|
||||
# Should now have the shell tool
|
||||
tool = manager.get_func("astrbot_execute_shell")
|
||||
assert tool is not None
|
||||
assert tool.name == "astrbot_execute_shell"
|
||||
@@ -40,7 +71,177 @@ def test_register_internal_tools_does_not_duplicate():
|
||||
first_tool = manager.get_func("astrbot_execute_shell")
|
||||
assert first_tool is not None
|
||||
|
||||
|
||||
# Should still have the same tool (not duplicated)
|
||||
second_tool = manager.get_func("astrbot_execute_shell")
|
||||
assert second_tool is first_tool
|
||||
assert first_tool is not None
|
||||
assert first_tool.parameters["properties"]["background"]["default"] is False # type: ignore[union-attr]
|
||||
assert manager.is_builtin_tool("astrbot_execute_shell") is True
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_execute_shell_defaults_to_foreground(monkeypatch):
|
||||
from astrbot.core.tools.computer_tools import shell as shell_tools
|
||||
|
||||
calls = []
|
||||
|
||||
class FakeShell:
|
||||
async def exec(self, command, cwd=None, background=False, env=None):
|
||||
calls.append({"command": command, "background": background})
|
||||
return {"success": True, "stdout": "", "stderr": "", "exit_code": 0}
|
||||
|
||||
class FakeBooter:
|
||||
shell = FakeShell()
|
||||
|
||||
FakeWrapper = _make_fake_wrapper_class()
|
||||
|
||||
async def fake_get_booter(context, session_id):
|
||||
return FakeBooter()
|
||||
|
||||
monkeypatch.setattr(shell_tools, "get_booter", fake_get_booter)
|
||||
|
||||
result = await ExecuteShellTool().call(
|
||||
FakeWrapper(), command="chromium https://example.com"
|
||||
)
|
||||
|
||||
assert isinstance(result, str)
|
||||
assert json.loads(result)["success"] is True
|
||||
assert calls == [{"command": "chromium https://example.com", "background": False}]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_execute_shell_uses_fresh_default_env_per_call(monkeypatch):
|
||||
from astrbot.core.tools.computer_tools import shell as shell_tools
|
||||
|
||||
calls = []
|
||||
|
||||
class FakeShell:
|
||||
async def exec(self, command, cwd=None, background=False, env=None):
|
||||
assert env is not None
|
||||
env["MUTATED_BY_FAKE_SHELL"] = command
|
||||
calls.append(env.copy())
|
||||
return {"success": True, "stdout": "", "stderr": "", "exit_code": 0}
|
||||
|
||||
class FakeBooter:
|
||||
shell = FakeShell()
|
||||
|
||||
FakeWrapper = _make_fake_wrapper_class()
|
||||
|
||||
async def fake_get_booter(context, session_id):
|
||||
return FakeBooter()
|
||||
|
||||
monkeypatch.setattr(shell_tools, "get_booter", fake_get_booter)
|
||||
tool = ExecuteShellTool()
|
||||
|
||||
await tool.call(FakeWrapper(), command="first")
|
||||
await tool.call(FakeWrapper(), command="second")
|
||||
|
||||
assert calls[0]["MUTATED_BY_FAKE_SHELL"] == "first"
|
||||
assert calls[1] == {"MUTATED_BY_FAKE_SHELL": "second"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_execute_shell_copies_user_env_before_execution(monkeypatch):
|
||||
from astrbot.core.tools.computer_tools import shell as shell_tools
|
||||
|
||||
calls = []
|
||||
|
||||
class FakeShell:
|
||||
async def exec(self, command, cwd=None, background=False, env=None):
|
||||
assert env is not None
|
||||
env["MUTATED_BY_FAKE_SHELL"] = command
|
||||
calls.append(env.copy())
|
||||
return {"success": True, "stdout": "", "stderr": "", "exit_code": 0}
|
||||
|
||||
class FakeBooter:
|
||||
shell = FakeShell()
|
||||
|
||||
FakeWrapper = _make_fake_wrapper_class()
|
||||
|
||||
async def fake_get_booter(context, session_id):
|
||||
return FakeBooter()
|
||||
|
||||
monkeypatch.setattr(shell_tools, "get_booter", fake_get_booter)
|
||||
original_env = {"FOO": "bar"}
|
||||
|
||||
await ExecuteShellTool().call(FakeWrapper(), command="first", env=original_env)
|
||||
|
||||
assert original_env == {"FOO": "bar"}
|
||||
assert calls == [{"FOO": "bar", "MUTATED_BY_FAKE_SHELL": "first"}]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_execute_shell_passes_background_flag_directly(monkeypatch):
|
||||
"""In the new architecture, background flag is passed directly to shell exec."""
|
||||
from astrbot.core.tools.computer_tools import shell as shell_tools
|
||||
|
||||
calls = []
|
||||
|
||||
class FakeShell:
|
||||
async def exec(self, command, cwd=None, background=False, env=None):
|
||||
calls.append({"command": command, "background": background})
|
||||
return {"success": True, "stdout": "", "stderr": "", "exit_code": 0}
|
||||
|
||||
class FakeBooter:
|
||||
shell = FakeShell()
|
||||
|
||||
FakeWrapper = _make_fake_wrapper_class()
|
||||
|
||||
async def fake_get_booter(context, session_id):
|
||||
return FakeBooter()
|
||||
|
||||
monkeypatch.setattr(shell_tools, "get_booter", fake_get_booter)
|
||||
|
||||
command = "nohup firefox >/tmp/astrbot-firefox.log 2>&1 &"
|
||||
result = await ExecuteShellTool().call(
|
||||
FakeWrapper(), command=command, background=True
|
||||
)
|
||||
|
||||
assert isinstance(result, str)
|
||||
assert json.loads(result)["success"] is True
|
||||
assert calls == [{"command": command, "background": True}]
|
||||
|
||||
command2 = "firefox & # already detached"
|
||||
result2 = await ExecuteShellTool().call(
|
||||
FakeWrapper(), command=command2, background=True
|
||||
)
|
||||
|
||||
assert isinstance(result2, str)
|
||||
assert json.loads(result2)["success"] is True
|
||||
assert calls[1] == {"command": command2, "background": True}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_execute_shell_reports_exception_type(monkeypatch):
|
||||
"""Error message uses e!s formatting (may omit class name if __str__ is blank)."""
|
||||
from astrbot.core.tools.computer_tools import shell as shell_tools
|
||||
|
||||
class FakeShell:
|
||||
async def exec(self, command, cwd=None, background=False, env=None):
|
||||
raise ValueError("custom error")
|
||||
|
||||
class FakeBooter:
|
||||
shell = FakeShell()
|
||||
|
||||
FakeWrapper = _make_fake_wrapper_class()
|
||||
|
||||
async def fake_get_booter(context, session_id):
|
||||
return FakeBooter()
|
||||
|
||||
monkeypatch.setattr(shell_tools, "get_booter", fake_get_booter)
|
||||
|
||||
result = await ExecuteShellTool().call(FakeWrapper(), command="firefox")
|
||||
|
||||
assert result == "Error executing command: custom error"
|
||||
|
||||
|
||||
def test_tavily_tools_are_registered_as_builtin_tools():
|
||||
manager = FunctionToolManager()
|
||||
|
||||
search_tool = manager.get_builtin_tool(TavilyWebSearchTool)
|
||||
extract_tool = manager.get_builtin_tool(TavilyExtractWebPageTool)
|
||||
|
||||
assert search_tool.name == "web_search_tavily"
|
||||
assert extract_tool.name == "tavily_extract_web_page"
|
||||
assert manager.is_builtin_tool("web_search_tavily") is True
|
||||
assert manager.is_builtin_tool("tavily_extract_web_page") is True
|
||||
|
||||
Reference in New Issue
Block a user