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:
LIghtJUNction
2026-04-29 02:06:24 +08:00
parent 4463895c89
commit 5f827fffd2
18 changed files with 648 additions and 362 deletions
+1 -1
View File
@@ -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)
-6
View File
@@ -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]
-22
View File
@@ -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})"
-1
View File
@@ -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]: ...
-115
View File
@@ -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
}
+5 -7
View File
@@ -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(
+21 -16
View File
@@ -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",
),
)
+101 -41
View File
@@ -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
-3
View File
@@ -4,11 +4,8 @@ from .python import PythonComponent
from .shell import ShellComponent
__all__ = [
"BrowserComponent",
"BrowserComponent",
"FileSystemComponent",
"FileSystemComponent",
"GUIComponent",
"PythonComponent",
"ShellComponent",
]
+1
View File
@@ -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"""
...
+251
View File
@@ -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 -2
View File
@@ -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
+15 -9
View File
@@ -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):
+28 -18
View File
@@ -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}"
+17 -8
View File
@@ -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"
+207 -6
View File
@@ -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