From 5f827fffd23991c9f43d40ef293cee3cbc5a2485 Mon Sep 17 00:00:00 2001 From: LIghtJUNction Date: Wed, 29 Apr 2026 02:06:24 +0800 Subject: [PATCH] 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. --- astrbot/core/agent/mcp_client.py | 2 +- astrbot/core/agent/run_context.py | 6 - astrbot/core/agent/tool.py | 22 -- astrbot/core/agent/tool_executor.py | 1 - astrbot/core/agent/tool_session_manager.py | 115 -------- astrbot/core/astr_agent_tool_exec.py | 12 +- astrbot/core/astr_main_agent.py | 37 +-- astrbot/core/computer/booters/local.py | 142 +++++++--- astrbot/core/computer/olayer/__init__.py | 3 - astrbot/core/computer/olayer/shell.py | 1 + astrbot/core/computer/shell_session.py | 251 ++++++++++++++++++ .../method/agent_sub_stages/third_party.py | 2 - astrbot/core/provider/sources/genie_tts.py | 3 +- astrbot/core/star/filter/command.py | 24 +- astrbot/core/tools/computer_tools/shell.py | 46 ++-- tests/benchmarks/test_performance.py | 25 +- .../test_agent/test_tool_session_manager.py | 105 -------- tests/unit/test_func_tool_manager.py | 213 ++++++++++++++- 18 files changed, 648 insertions(+), 362 deletions(-) delete mode 100644 astrbot/core/agent/tool_session_manager.py create mode 100644 astrbot/core/computer/shell_session.py delete mode 100644 tests/unit/test_core/test_agent/test_tool_session_manager.py diff --git a/astrbot/core/agent/mcp_client.py b/astrbot/core/agent/mcp_client.py index 5a6d9a56d..080758c7b 100644 --- a/astrbot/core/agent/mcp_client.py +++ b/astrbot/core/agent/mcp_client.py @@ -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) diff --git a/astrbot/core/agent/run_context.py b/astrbot/core/agent/run_context.py index dfdf256c1..3c500b2d6 100644 --- a/astrbot/core/agent/run_context.py +++ b/astrbot/core/agent/run_context.py @@ -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] diff --git a/astrbot/core/agent/tool.py b/astrbot/core/agent/tool.py index 227b95db5..41581ead6 100644 --- a/astrbot/core/agent/tool.py +++ b/astrbot/core/agent/tool.py @@ -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})" diff --git a/astrbot/core/agent/tool_executor.py b/astrbot/core/agent/tool_executor.py index 0ba2c9522..14fe4beee 100644 --- a/astrbot/core/agent/tool_executor.py +++ b/astrbot/core/agent/tool_executor.py @@ -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]: ... diff --git a/astrbot/core/agent/tool_session_manager.py b/astrbot/core/agent/tool_session_manager.py deleted file mode 100644 index 686fa21f0..000000000 --- a/astrbot/core/agent/tool_session_manager.py +++ /dev/null @@ -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 - } diff --git a/astrbot/core/astr_agent_tool_exec.py b/astrbot/core/astr_agent_tool_exec.py index be5b166ca..1ac4479f3 100644 --- a/astrbot/core/astr_agent_tool_exec.py +++ b/astrbot/core/astr_agent_tool_exec.py @@ -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( diff --git a/astrbot/core/astr_main_agent.py b/astrbot/core/astr_main_agent.py index 9bde20bce..39896830a 100644 --- a/astrbot/core/astr_main_agent.py +++ b/astrbot/core/astr_main_agent.py @@ -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"\n{quoted_content}\n" @@ -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", ), ) diff --git a/astrbot/core/computer/booters/local.py b/astrbot/core/computer/booters/local.py index 41ac40657..4d7a29d96 100644 --- a/astrbot/core/computer/booters/local.py +++ b/astrbot/core/computer/booters/local.py @@ -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 diff --git a/astrbot/core/computer/olayer/__init__.py b/astrbot/core/computer/olayer/__init__.py index dd73e9e90..261f9de9c 100644 --- a/astrbot/core/computer/olayer/__init__.py +++ b/astrbot/core/computer/olayer/__init__.py @@ -4,11 +4,8 @@ from .python import PythonComponent from .shell import ShellComponent __all__ = [ - "BrowserComponent", "BrowserComponent", "FileSystemComponent", - "FileSystemComponent", - "GUIComponent", "PythonComponent", "ShellComponent", ] diff --git a/astrbot/core/computer/olayer/shell.py b/astrbot/core/computer/olayer/shell.py index fb9ccdcae..a563da603 100644 --- a/astrbot/core/computer/olayer/shell.py +++ b/astrbot/core/computer/olayer/shell.py @@ -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""" ... diff --git a/astrbot/core/computer/shell_session.py b/astrbot/core/computer/shell_session.py new file mode 100644 index 000000000..02326a66c --- /dev/null +++ b/astrbot/core/computer/shell_session.py @@ -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) diff --git a/astrbot/core/pipeline/process_stage/method/agent_sub_stages/third_party.py b/astrbot/core/pipeline/process_stage/method/agent_sub_stages/third_party.py index 3b5b2c235..5602f9612 100644 --- a/astrbot/core/pipeline/process_stage/method/agent_sub_stages/third_party.py +++ b/astrbot/core/pipeline/process_stage/method/agent_sub_stages/third_party.py @@ -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, diff --git a/astrbot/core/provider/sources/genie_tts.py b/astrbot/core/provider/sources/genie_tts.py index aeb3a1581..3a4cc125c 100644 --- a/astrbot/core/provider/sources/genie_tts.py +++ b/astrbot/core/provider/sources/genie_tts.py @@ -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 diff --git a/astrbot/core/star/filter/command.py b/astrbot/core/star/filter/command.py index b9b36d5ac..af07100e8 100644 --- a/astrbot/core/star/filter/command.py +++ b/astrbot/core/star/filter/command.py @@ -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): diff --git a/astrbot/core/tools/computer_tools/shell.py b/astrbot/core/tools/computer_tools/shell.py index 538ebe397..dac9333dc 100644 --- a/astrbot/core/tools/computer_tools/shell.py +++ b/astrbot/core/tools/computer_tools/shell.py @@ -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}" diff --git a/tests/benchmarks/test_performance.py b/tests/benchmarks/test_performance.py index 53900169a..f70666e75 100644 --- a/tests/benchmarks/test_performance.py +++ b/tests/benchmarks/test_performance.py @@ -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 diff --git a/tests/unit/test_core/test_agent/test_tool_session_manager.py b/tests/unit/test_core/test_agent/test_tool_session_manager.py deleted file mode 100644 index 14fc5db81..000000000 --- a/tests/unit/test_core/test_agent/test_tool_session_manager.py +++ /dev/null @@ -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" diff --git a/tests/unit/test_func_tool_manager.py b/tests/unit/test_func_tool_manager.py index 244fc5af4..33801ba80 100644 --- a/tests/unit/test_func_tool_manager.py +++ b/tests/unit/test_func_tool_manager.py @@ -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