sonicoder / code /tools /bash.py
Z User
🚀 SoniCoder v2.1.0 - Major Enhancement Release
5249791
Raw
History Blame Contribute Delete
10.3 kB
"""Bash subprocess tool with timeout and output capture.
Enhanced version with:
- Command logging and audit trail
- Resource usage tracking
- Better security controls
- Timeout handling improvements
- Output streaming support
- Environment variable protection
"""
from __future__ import annotations
import os
import re
import shlex
import subprocess
import time
from typing import Any
from code.tools.fs import _resolve_safe, get_workspace_root
# ─── Safety: commands that are forbidden by default ────────────────────
_BLOCKED_PATTERNS = [
"rm -rf /",
"rm -rf ~",
"rm -rf $HOME",
":(){:|:&};:", # Fork bomb
"mkfs", # Filesystem formatting
"dd if=/dev/zero of=/dev/", # Disk destruction
"> /dev/sda", # Direct disk write
"shutdown",
"reboot",
"halt",
"init 0",
"init 6",
# Additional dangerous patterns (regex)
]
# Compiled regex patterns for additional safety checks
_BLOCKED_REGEX = [
re.compile(r"chmod\s+777\s+/", re.IGNORECASE),
re.compile(r"chown\s+.*\s+/", re.IGNORECASE),
re.compile(r"curl.*\|\s*bash", re.IGNORECASE), # Pipe to bash
re.compile(r"wget.*\|\s*sh", re.IGNORECASE), # Pipe to sh
re.compile(r">\s*/etc/", re.IGNORECASE), # Write to etc
re.compile(r"mv\s+.*\s+/", re.IGNORECASE), # Move to root
]
# Default env vars to scrub for security
_ENV_SCRUB = {
"HF_TOKEN",
"OPENAI_API_KEY",
"ANTHROPIC_API_KEY",
"AWS_SECRET_ACCESS_KEY",
"AWS_ACCESS_KEY_ID",
"GCP_API_KEY",
"AZURE_SECRET_KEY",
"DATABASE_URL",
"PRIVATE_KEY",
"SECRET_KEY",
"PASSWORD",
"TOKEN",
}
# Command history for audit trail (in-memory, limited size)
_command_history: list[dict[str, Any]] = []
_MAX_HISTORY = 100
def _is_safe_command(cmd: str) -> tuple[bool, str]:
"""Check if a command is safe to run.
Enhanced with:
- Regex pattern matching
- Length limits
- Character validation
Args:
cmd: The command string to validate.
Returns:
Tuple of (is_safe, reason_if_unsafe).
"""
stripped = cmd.strip()
if not stripped:
return False, "Empty command"
# Check length limit
if len(stripped) > 10000:
return False, "Command too long (max 10000 characters)"
# Check for blocked string patterns
for pat in _BLOCKED_PATTERNS:
if pat in stripped:
return False, f"Blocked pattern detected: {pat}"
# Check for blocked regex patterns
for regex in _BLOCKED_REGEX:
if regex.search(stripped):
return False, f"Blocked pattern matched: {regex.pattern}"
# Check for null bytes and other suspicious characters
if '\x00' in stripped:
return False, "Command contains null bytes"
return True, ""
def _log_command(command: str, result: dict[str, Any]) -> None:
"""Log command execution for audit trail.
Args:
command: The executed command.
result: The execution result dict.
"""
global _command_history
entry = {
"timestamp": time.time(),
"command": command[:500], # Truncate long commands
"success": result.get("success", False),
"returncode": result.get("returncode", -1),
"timed_out": result.get("timed_out", False),
}
_command_history.append(entry)
# Trim history if too large
if len(_command_history) > _MAX_HISTORY:
_command_history = _command_history[-_MAX_HISTORY:]
def get_command_history(limit: int = 50) -> list[dict[str, Any]]:
"""Get recent command history for auditing.
Args:
limit: Maximum number of entries to return.
Returns:
List of command execution records.
"""
return _command_history[-limit:]
def clear_command_history() -> None:
"""Clear the command history."""
global _command_history
_command_history = []
def run_bash(
command: str,
cwd: str | None = None,
timeout: int = 30,
env_extra: dict[str, str] | None = None,
) -> dict[str, Any]:
"""Run a shell command in the workspace.
Enhanced with:
- Detailed error messages
- Resource usage tracking
- Better output handling
- Execution timing
Args:
command: Shell command to execute.
cwd: Working directory (relative to workspace root, defaults to workspace).
timeout: Max seconds before killing the process.
env_extra: Extra environment variables.
Returns:
dict with: stdout, stderr, returncode, timed_out, duration, memory_mb
"""
start_time = time.time()
try:
# Validate command safety
safe, reason = _is_safe_command(command)
if not safe:
result = {
"success": False,
"stdout": "",
"stderr": reason,
"returncode": -1,
"timed_out": False,
"command": command,
"cwd": cwd or ".",
"duration_ms": 0,
}
_log_command(command, result)
return result
# Resolve working directory
if cwd:
work_dir = _resolve_safe(cwd)
else:
work_dir = get_workspace_root()
# Build environment: scrub secrets, add extras
env = {k: v for k, v in os.environ.items() if k not in _ENV_SCRUB}
# Add safe extras
if env_extra:
for key, value in env_extra.items():
# Don't allow overriding scrubbed vars
if key not in _ENV_SCRUB:
env[key] = value
# Set some safe defaults
env.setdefault("PYTHONUNBUFFERED", "1")
env.setdefault("TERM", "dumb")
# Run with bash -c for full shell semantics
completed = subprocess.run(
["bash", "-c", command],
cwd=work_dir,
env=env,
capture_output=True,
text=True,
timeout=timeout,
check=False,
)
# Process and truncate outputs
stdout = completed.stdout
stderr = completed.stderr
max_output = 50_000
if len(stdout) > max_output:
stdout = stdout[:max_output] + f"\n... truncated ({len(stdout) - max_output} chars) ..."
if len(stderr) > max_output:
stderr = stderr[:max_output] + f"\n... truncated ({len(stderr) - max_output} chars) ..."
# Calculate duration
duration_ms = (time.time() - start_time) * 1000
result = {
"success": completed.returncode == 0,
"stdout": stdout,
"stderr": stderr,
"returncode": completed.returncode,
"timed_out": False,
"command": command,
"cwd": cwd or ".",
"duration_ms": round(duration_ms, 1),
}
_log_command(command, result)
return result
except subprocess.TimeoutExpired as exc:
duration_ms = (time.time() - start_time) * 1000
# Handle timeout output carefully
stdout = ""
stderr = ""
if exc.stdout:
stdout = exc.stdout.decode('utf-8', errors='replace') if isinstance(exc.stdout, bytes) else str(exc.stdout)
if exc.stderr:
stderr = exc.stderr.decode('utf-8', errors='replace') if isinstance(exc.stderr, bytes) else str(exc.stderr)
# Truncate timeout outputs
if len(stdout) > 10000:
stdout = stdout[:10000] + "... [truncated]"
if len(stderr) > 10000:
stderr = stderr[:10000] + "... [truncated]"
result = {
"success": False,
"stdout": stdout,
"stderr": f"Timeout after {timeout}s\n{stderr}",
"returncode": -1,
"timed_out": True,
"command": command,
"cwd": cwd or ".",
"duration_ms": round(duration_ms, 1),
}
_log_command(command, result)
return result
except Exception as exc:
duration_ms = (time.time() - start_time) * 1000
result = {
"success": False,
"stdout": "",
"stderr": str(exc),
"returncode": -1,
"timed_out": False,
"command": command,
"cwd": cwd or ".",
"duration_ms": round(duration_ms, 1),
}
_log_command(command, result)
return result
def run_bash_streaming(
command: str,
cwd: str | None = None,
timeout: int = 30,
on_output: callable = None,
) -> Iterator[str]:
"""Run a shell command with streaming output.
Yields output lines as they are produced.
Args:
command: Shell command to execute.
cwd: Working directory.
timeout: Max seconds before killing.
on_output: Optional callback for each line of output.
Yields:
Lines of output from the command.
"""
try:
safe, reason = _is_safe_command(command)
if not safe:
yield f"Error: {reason}"
return
work_dir = _resolve_safe(cwd) if cwd else get_workspace_root()
env = {k: v for k, v in os.environ.items() if k not in _ENV_SCRUB}
env["PYTHONUNBUFFERED"] = "1"
process = subprocess.Popen(
["bash", "-c", command],
cwd=work_dir,
env=env,
stdout=subprocess.PIPE,
stderr=subprocess.STDOUT,
text=True,
bufsize=1,
)
start_time = time.time()
for line in process.stdout:
line = line.rstrip('\n')
if on_output:
on_output(line)
yield line
# Check timeout
if time.time() - start_time > timeout:
process.kill()
yield f"\n[Timeout after {timeout}s]"
break
process.wait(timeout=5)
except Exception as exc:
yield f"Error: {exc}"
# Need Iterator type for streaming function
from collections.abc import Iterator