ForcePilot/backend/package/yuxi/agents/backends/sandbox/backend.py
2026-06-02 18:47:28 +08:00

540 lines
22 KiB
Python

from __future__ import annotations
import base64
from datetime import datetime
from pathlib import PurePosixPath
from typing import Any
from deepagents.backends.protocol import (
EditResult,
ExecuteResponse,
FileDownloadResponse,
FileInfo,
FileUploadResponse,
GlobResult,
GrepMatch,
GrepResult,
LsResult,
ReadResult,
WriteResult,
)
from deepagents.backends.sandbox import BaseSandbox
from yuxi import config as conf
from yuxi.agents.skills.service import sync_thread_readable_skills
from yuxi.utils.logging_config import logger
from yuxi.utils.paths import (
OUTPUTS_DIR_NAME,
UPLOADS_DIR_NAME,
VIRTUAL_PATH_PREFIX,
VIRTUAL_SKILLS_PATH,
WORKSPACE_DIR_NAME,
)
from .provider import get_sandbox_provider, sandbox_id_for_thread
_USER_DATA_ROOT = "/" + VIRTUAL_PATH_PREFIX.strip("/")
_WORKSPACE_ROOT = f"{_USER_DATA_ROOT}/{WORKSPACE_DIR_NAME}"
_UPLOADS_ROOT = f"{_USER_DATA_ROOT}/{UPLOADS_DIR_NAME}"
_OUTPUTS_ROOT = f"{_USER_DATA_ROOT}/{OUTPUTS_DIR_NAME}"
_SKILLS_ROOT = "/" + VIRTUAL_SKILLS_PATH.strip("/")
_READABLE_ROOTS = (_USER_DATA_ROOT, _SKILLS_ROOT)
_WRITABLE_ROOTS = (_WORKSPACE_ROOT, _OUTPUTS_ROOT)
def _normalize_path(path: str) -> str:
raw = str(path or "").strip()
if not raw:
raise ValueError("path is required")
if not raw.startswith("/"):
raise ValueError("path must start with /")
pure = PurePosixPath(raw)
if ".." in pure.parts:
raise ValueError("path traversal is not allowed")
return str(pure)
def _is_same_or_child(path: str, root: str) -> bool:
root = root.rstrip("/") or "/"
if root == "/":
return path == "/" or path.startswith("/")
return path == root or path.startswith(f"{root}/")
def _path_overlaps_root(path: str, root: str) -> bool:
return _is_same_or_child(path, root) or _is_same_or_child(root, path)
def _can_read_path(path: str) -> bool:
return any(_is_same_or_child(path, root) for root in _READABLE_ROOTS)
def _can_list_path(path: str) -> bool:
return any(_path_overlaps_root(path, root) for root in _READABLE_ROOTS)
def _can_write_path(path: str) -> bool:
return any(_is_same_or_child(path, root) for root in _WRITABLE_ROOTS)
def _readable_search_paths(path: str) -> list[str]:
if _can_read_path(path):
return [path]
return [root for root in _READABLE_ROOTS if _is_same_or_child(root, path)]
def _glob_for_search_root(pattern: str, root: str) -> str:
bare_pattern = str(pattern or "*").lstrip("/")
bare_root = root.strip("/")
if bare_pattern == bare_root:
return "*"
root_prefix = f"{bare_root}/"
if bare_pattern.startswith(root_prefix):
return bare_pattern[len(root_prefix) :] or "*"
return pattern
def _filter_readable_infos(infos: list[FileInfo]) -> list[FileInfo]:
result: list[FileInfo] = []
for info in infos:
try:
path = _normalize_path(info.get("path", ""))
except ValueError:
continue
if _can_list_path(path):
result.append(info)
return result
def _filter_readable_matches(matches: list[GrepMatch]) -> list[GrepMatch]:
result: list[GrepMatch] = []
for match in matches:
try:
path = _normalize_path(match.get("path", ""))
except ValueError:
continue
if _can_read_path(path):
result.append(match)
return result
def _permission_error(operation: str, path: str) -> str:
return f"permission denied for {operation} on '{path}'"
def _describe_read_error(file_path: str, exc: Exception) -> str:
if isinstance(exc, FileNotFoundError):
return f"Error: File '{file_path}' not found"
if isinstance(exc, IsADirectoryError):
return f"Error: Path '{file_path}' is a directory"
if isinstance(exc, PermissionError):
return f"Error: Access denied for '{file_path}'"
if isinstance(exc, ValueError):
return f"Error: Invalid path '{file_path}': {exc}"
detail = str(exc).strip()
if detail:
return f"Error: Failed to read '{file_path}': {detail}"
return f"Error: Failed to read '{file_path}'"
def _looks_like_binary(content: bytes) -> bool:
if not content:
return False
if b"\x00" in content:
return True
try:
content.decode("utf-8")
return False
except UnicodeDecodeError:
return True
class ProvisionerSandboxBackend(BaseSandbox):
def __init__(
self,
thread_id: str,
*,
uid: str,
readable_skills: list[str] | None = None,
file_thread_id: str | None = None,
skills_thread_id: str | None = None,
):
self._thread_id = str(thread_id or "").strip()
if not self._thread_id:
raise ValueError("thread_id is required for ProvisionerSandboxBackend")
self._file_thread_id = str(file_thread_id or self._thread_id).strip()
if not self._file_thread_id:
raise ValueError("file_thread_id is required for ProvisionerSandboxBackend")
self._skills_thread_id = str(skills_thread_id or self._thread_id).strip()
if not self._skills_thread_id:
raise ValueError("skills_thread_id is required for ProvisionerSandboxBackend")
self._uid = str(uid or "").strip()
if not self._uid:
raise ValueError("uid is required for ProvisionerSandboxBackend")
self._readable_skills = list(readable_skills or [])
self._provider = get_sandbox_provider()
self._id = sandbox_id_for_thread(self._file_thread_id, self._skills_thread_id, uid=self._uid)
self._client: Any | None = None
self._client_url: str | None = None
self._command_timeout_seconds = int(getattr(conf, "sandbox_exec_timeout_seconds", 180))
self._max_output_bytes = int(getattr(conf, "sandbox_max_output_bytes", 262_144))
@property
def id(self) -> str:
return self._id
def _build_client(self, sandbox_url: str):
try:
from agent_sandbox import Sandbox as AgentSandboxClient
except Exception as exc: # noqa: BLE001
raise RuntimeError(
"agent-sandbox is required. Install dependency `agent-sandbox` in the docker image."
) from exc
return AgentSandboxClient(base_url=sandbox_url, timeout=self._command_timeout_seconds)
def _get_client(self) -> Any:
sync_thread_readable_skills(self._skills_thread_id, self._readable_skills)
connection = self._provider.get(
self._thread_id,
uid=self._uid,
create_if_missing=True,
file_thread_id=self._file_thread_id,
skills_thread_id=self._skills_thread_id,
)
if connection is None:
raise RuntimeError(f"sandbox is unavailable for thread {self._thread_id}")
if self._client is None or self._client_url != connection.sandbox_url:
self._client = self._build_client(connection.sandbox_url)
self._client_url = connection.sandbox_url
return self._client
def _read_binary(self, path: str, offset: int = 0, limit: int | None = None) -> bytes:
"""Read file content from the sandbox file API and normalize it to bytes.
The underlying API returns plain text by default and may include an
explicit `encoding="base64"` marker for binary payloads. This helper is
the single normalization point used by read(), edit(), and download_files().
"""
start_line = max(0, int(offset))
end_line = start_line + int(limit) if limit is not None else None
result = self._get_client().file.read_file(
file=path,
start_line=start_line,
end_line=end_line,
)
content = result.data.content
if content is None:
return b""
if isinstance(content, bytes):
return content
if not isinstance(content, str):
return str(content).encode("utf-8")
encoding = getattr(result.data, "encoding", None)
if isinstance(encoding, str) and encoding.lower() == "base64":
return base64.b64decode(content, validate=True)
return content.encode("utf-8")
def read(
self,
file_path: str,
offset: int = 0,
limit: int = 2000,
) -> ReadResult:
"""Read allowed file content via the sandbox file API."""
try:
normalized_path = _normalize_path(file_path)
except Exception as exc: # noqa: BLE001
return ReadResult(error=f"Invalid path '{file_path}': {exc}")
if not _can_read_path(normalized_path):
return ReadResult(error=_permission_error("read", normalized_path))
try:
content = self._read_binary(normalized_path, offset=offset, limit=limit)
if _looks_like_binary(content):
content = self._read_binary(normalized_path)
except Exception as exc: # noqa: BLE001
error = _describe_read_error(file_path, exc)
return ReadResult(error=error.removeprefix("Error: "))
if _looks_like_binary(content):
return ReadResult(file_data={"content": base64.b64encode(content).decode("ascii"), "encoding": "base64"})
return ReadResult(file_data={"content": content.decode("utf-8"), "encoding": "utf-8"})
def execute(self, command: str, *, timeout: int | None = None) -> ExecuteResponse:
"""Execute a shell command in the sandbox.
Output is normalized to text and truncated to the configured maximum
payload size before being returned.
"""
try:
kwargs: dict[str, Any] = {"command": command}
if timeout is not None:
kwargs["timeout"] = timeout
result = self._get_client().shell.exec_command(**kwargs)
output = result.data.output or ""
exit_code = result.data.exit_code
truncated = False
encoded = output.encode("utf-8", errors="ignore")
if len(encoded) > self._max_output_bytes:
output = encoded[: self._max_output_bytes].decode("utf-8", errors="ignore")
truncated = True
return ExecuteResponse(
output=output,
exit_code=exit_code if isinstance(exit_code, int) else None,
truncated=truncated,
)
except Exception as exc: # noqa: BLE001
logger.error(f"Sandbox execute failed for thread {self._thread_id}: {exc}")
return ExecuteResponse(output=f"Error: {exc}", exit_code=1, truncated=False)
def ls(self, path: str) -> LsResult:
"""List direct children of an allowed sandbox path with lightweight metadata."""
try:
normalized_path = _normalize_path(path)
except Exception as exc: # noqa: BLE001
return LsResult(error=f"Invalid path '{path}': {exc}")
if not _can_list_path(normalized_path):
return LsResult(error=_permission_error("read", normalized_path))
try:
result = self._get_client().file.list_path(path=normalized_path, recursive=False, include_size=True)
except Exception as exc: # noqa: BLE001
return LsResult(error=str(exc) or f"Failed to list '{path}'")
entries = result.data.files or []
infos: list[FileInfo] = []
for entry in entries:
info: FileInfo = {"path": entry.path, "is_dir": entry.is_directory}
size = entry.size
if isinstance(size, int):
info["size"] = size
modified_time = entry.modified_time
if modified_time:
if isinstance(modified_time, str) and modified_time.isdigit():
info["modified_at"] = datetime.fromtimestamp(int(modified_time)).isoformat()
elif isinstance(modified_time, str):
try:
info["modified_at"] = datetime.fromisoformat(modified_time).isoformat()
except ValueError:
info["modified_at"] = modified_time
elif isinstance(modified_time, (int, float)):
info["modified_at"] = datetime.fromtimestamp(modified_time).isoformat()
infos.append(info)
return LsResult(entries=_filter_readable_infos(infos))
def write(self, file_path: str, content: str) -> WriteResult:
"""Write a new text file.
This method is intentionally text-only. Binary payloads should go through
upload_files(), which uses base64 encoding for the sandbox file API.
"""
try:
normalized_path = _normalize_path(file_path)
except Exception as exc: # noqa: BLE001
return WriteResult(error=f"Error: Invalid path '{file_path}': {exc}")
if not _can_write_path(normalized_path):
return WriteResult(error=f"Error: {_permission_error('write', normalized_path)}")
if not isinstance(content, str):
return WriteResult(error="Error: write() only supports text content; use upload_files() for binary data")
try:
self._read_binary(normalized_path)
except Exception: # noqa: BLE001
pass
else:
return WriteResult(error=f"Error: File '{file_path}' already exists")
try:
result = self._get_client().file.write_file(file=normalized_path, content=content)
if not result.success:
return WriteResult(error=result.message or f"Failed to write file '{file_path}'")
except Exception as exc: # noqa: BLE001
return WriteResult(error=str(exc) or f"Failed to write file '{file_path}'")
return WriteResult(path=normalized_path)
def edit(
self,
file_path: str,
old_string: str,
new_string: str,
replace_all: bool = False, # noqa: FBT001, FBT002
) -> EditResult:
"""Edit an existing text file by replacing string content.
This method operates on UTF-8-decoded text content only. Binary files
are not supported here and should be handled via download/upload flows.
"""
try:
normalized_path = _normalize_path(file_path)
except Exception as exc: # noqa: BLE001
return EditResult(error=f"Error: Invalid path '{file_path}': {exc}")
if not _can_write_path(normalized_path):
return EditResult(error=f"Error: {_permission_error('write', normalized_path)}")
# Check if old_string exists
try:
text = self._read_binary(normalized_path).decode("utf-8", errors="replace")
except Exception: # noqa: BLE001
return EditResult(error=f"Error: File '{file_path}' not found")
count = text.count(old_string)
if count == 0:
return EditResult(error=f"Error: String not found in file: '{old_string}'")
if count > 1 and not replace_all:
return EditResult(
error=(
f"Error: String '{old_string}' appears multiple times. "
"Use replace_all=True to replace all occurrences."
)
)
# Use str_replace_editor API
replace_mode = "ALL" if replace_all else "FIRST"
try:
result = self._get_client().file.str_replace_editor(
command="str_replace",
path=normalized_path,
old_str=old_string,
new_str=new_string,
replace_mode=replace_mode,
)
if not result.success:
return EditResult(error=result.message or f"Error editing file '{file_path}'")
except Exception as exc: # noqa: BLE001
return EditResult(error=f"Error editing file: {exc}")
return EditResult(path=normalized_path, occurrences=count if replace_all else 1)
def grep(
self,
pattern: str,
path: str | None = None,
glob: str | None = None,
) -> GrepResult:
"""Search allowed sandbox paths for literal text."""
try:
normalized_path = _normalize_path(path or "/")
except Exception as exc: # noqa: BLE001
return GrepResult(error=f"Invalid path '{path or '/'}': {exc}")
search_paths = _readable_search_paths(normalized_path)
if not search_paths:
return GrepResult(error=_permission_error("read", normalized_path))
matches: list[GrepMatch] = []
for search_path in search_paths:
result = super().grep(pattern=pattern, path=search_path, glob=glob)
if result.error:
return result
matches.extend(result.matches or [])
return GrepResult(matches=_filter_readable_matches(matches))
def glob(self, pattern: str, path: str = "/") -> GlobResult:
"""Return files matching a glob pattern under allowed sandbox paths."""
try:
normalized_path = _normalize_path(path)
except Exception as exc: # noqa: BLE001
return GlobResult(error=f"Invalid path '{path}': {exc}")
if ".." in PurePosixPath(str(pattern or "")).parts:
return GlobResult(error="Invalid glob pattern: path traversal is not allowed")
search_paths = _readable_search_paths(normalized_path)
if not search_paths:
return GlobResult(error=_permission_error("read", normalized_path))
infos: list[FileInfo] = []
for search_path in search_paths:
try:
result = self._get_client().file.find_files(
path=search_path,
glob=_glob_for_search_root(pattern, search_path),
)
except Exception as exc: # noqa: BLE001
return GlobResult(error=str(exc) or f"Failed to glob '{path}'")
for file_path in result.data.files or []:
infos.append({"path": file_path})
infos = _filter_readable_infos(infos)
infos.sort(key=lambda item: item.get("path", ""))
return GlobResult(matches=infos)
def upload_files(self, files: list[tuple[str, bytes]]) -> list[FileUploadResponse]:
"""Upload binary or text file payloads via the sandbox file API.
Contents are base64-encoded before calling the remote write_file API so
arbitrary bytes can be transferred safely.
"""
responses: list[FileUploadResponse] = []
for path, content in files:
try:
normalized_path = _normalize_path(path)
if not _can_write_path(normalized_path):
responses.append(FileUploadResponse(path=normalized_path, error="permission_denied"))
continue
result = self._get_client().file.write_file(
file=normalized_path,
content=base64.b64encode(content).decode("ascii"),
encoding="base64",
)
if not result.success:
raise Exception(result.message or "Upload failed")
responses.append(FileUploadResponse(path=normalized_path, error=None))
except PermissionError:
normalized_path = str(path)
responses.append(FileUploadResponse(path=normalized_path, error="permission_denied"))
except IsADirectoryError:
normalized_path = str(path)
responses.append(FileUploadResponse(path=normalized_path, error="is_directory"))
except FileNotFoundError:
normalized_path = str(path)
responses.append(FileUploadResponse(path=normalized_path, error="file_not_found"))
except Exception as exc: # noqa: BLE001
normalized_path = str(path)
logger.warning(f"Upload to sandbox failed for {normalized_path}: {exc}")
responses.append(FileUploadResponse(path=normalized_path, error="invalid_path"))
return responses
def download_files(self, paths: list[str]) -> list[FileDownloadResponse]:
"""Download file payloads as raw bytes from the sandbox file API.
_read_binary() normalizes the sandbox file API response to bytes.
"""
responses: list[FileDownloadResponse] = []
for path in paths:
try:
normalized_path = _normalize_path(path)
if not _can_read_path(normalized_path):
responses.append(
FileDownloadResponse(path=normalized_path, content=None, error="permission_denied")
)
continue
content = self._read_binary(normalized_path)
responses.append(FileDownloadResponse(path=normalized_path, content=content, error=None))
except PermissionError:
normalized_path = str(path)
responses.append(FileDownloadResponse(path=normalized_path, content=None, error="permission_denied"))
except IsADirectoryError:
normalized_path = str(path)
responses.append(FileDownloadResponse(path=normalized_path, content=None, error="is_directory"))
except FileNotFoundError:
normalized_path = str(path)
responses.append(FileDownloadResponse(path=normalized_path, content=None, error="file_not_found"))
except ValueError:
normalized_path = str(path)
responses.append(FileDownloadResponse(path=normalized_path, content=None, error="invalid_path"))
except Exception as exc: # noqa: BLE001
normalized_path = str(path)
logger.warning(f"Download from sandbox failed for {normalized_path}: {exc}")
responses.append(FileDownloadResponse(path=normalized_path, content=None, error=f"read_failed: {exc}"))
return responses